Initial commit: TaskPool React panel
- React frontend with route-level code splitting - Backend rebranded from Baihu to TaskPool - DB brand migration script and local compatibility
This commit is contained in:
@@ -0,0 +1,169 @@
|
||||
package bootstrap
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/executor"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/router"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/tunnel"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type App struct {
|
||||
Config *services.AppConfig
|
||||
Router *gin.Engine
|
||||
}
|
||||
|
||||
func New() *App {
|
||||
app := InitBasic()
|
||||
app.initRouter()
|
||||
|
||||
// 初始化完成后将路由引擎注入到隧道模块,以支持高性能的纯内存代理
|
||||
tunnel.SetLocalEngine(app.Router)
|
||||
|
||||
// 初始化隧道后台服务 (读取配置决定角色并启动服务)
|
||||
tunnel.Init()
|
||||
|
||||
// 启动系统级后台定时任务调度器
|
||||
executor.InitSysCron()
|
||||
|
||||
// 初始化完成后回收一次内存
|
||||
utils.FreeMemory()
|
||||
|
||||
return app
|
||||
}
|
||||
|
||||
var (
|
||||
globalApp *App
|
||||
initOnce sync.Once
|
||||
)
|
||||
|
||||
func InitBasic() *App {
|
||||
initOnce.Do(func() {
|
||||
app := &App{}
|
||||
utils.InitRuntime()
|
||||
utils.InitSecretKey()
|
||||
|
||||
// 自动加载配置 (内部会自动处理 BH_CONFIG_PATH 环境变量与默认路径的优先级)
|
||||
app.initConfigWithPath("")
|
||||
app.initDatabase()
|
||||
|
||||
logger.Infof("[System] 低于1.0.11版本升级最新版本错误指引: https://github.com/engigu/taskpool/issues/64")
|
||||
globalApp = app
|
||||
})
|
||||
return globalApp
|
||||
}
|
||||
|
||||
// InitBasicForCmd 专为命令行工具定制的基础环境初始化入口
|
||||
// 内部会调高控制台日志过滤级别以自动静默屏蔽刷屏的底层系统与组件启动 Info 日志
|
||||
func InitBasicForCmd() *App {
|
||||
logger.SetLevel("warn")
|
||||
return InitBasic()
|
||||
}
|
||||
|
||||
// func (a *App) initConfig() {
|
||||
// a.initConfigWithPath(constant.ConfigPath)
|
||||
// }
|
||||
|
||||
func (a *App) initConfigWithPath(path string) {
|
||||
cfg, err := services.LoadConfig(path)
|
||||
if err != nil {
|
||||
logger.Fatalf("Failed to load config: %v", err)
|
||||
}
|
||||
a.Config = cfg
|
||||
|
||||
// Ensure directories exist
|
||||
err = os.MkdirAll(constant.DataDir, 0755)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
err = os.MkdirAll(constant.ScriptsWorkDir, 0755)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
a.setupTaskPoolBin()
|
||||
}
|
||||
|
||||
func (a *App) setupTaskPoolBin() {
|
||||
binDir := filepath.Join(constant.DataDir, "bin")
|
||||
_ = os.MkdirAll(binDir, 0755)
|
||||
|
||||
exe, err := os.Executable()
|
||||
if err == nil {
|
||||
// 新命令名
|
||||
linkPath := filepath.Join(binDir, "taskpool")
|
||||
if runtime.GOOS == "windows" {
|
||||
linkPath += ".exe"
|
||||
}
|
||||
os.Remove(linkPath)
|
||||
_ = os.Symlink(exe, linkPath)
|
||||
|
||||
// 兼容旧命令名 baihu
|
||||
legacyPath := filepath.Join(binDir, "baihu")
|
||||
if runtime.GOOS == "windows" {
|
||||
legacyPath += ".exe"
|
||||
}
|
||||
os.Remove(legacyPath)
|
||||
_ = os.Symlink(exe, legacyPath)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) initDatabase() {
|
||||
dbCfg := &database.Config{
|
||||
Type: a.Config.Database.Type,
|
||||
Host: a.Config.Database.Host,
|
||||
Port: a.Config.Database.Port,
|
||||
User: a.Config.Database.User,
|
||||
Password: a.Config.Database.Password,
|
||||
DBName: a.Config.Database.DBName,
|
||||
Path: a.Config.Database.Path,
|
||||
DSN: a.Config.Database.DSN,
|
||||
SSLMode: a.Config.Database.SSLMode,
|
||||
}
|
||||
|
||||
if err := database.Init(dbCfg); err != nil {
|
||||
logger.Fatalf("Failed to init database: %v", err)
|
||||
}
|
||||
|
||||
// 记录各个初始化阶段的时间
|
||||
startTime := time.Now()
|
||||
|
||||
// 执行 V3 迁移(ID 变更迁移)
|
||||
if err := services.RunMigrationV3(); err != nil {
|
||||
logger.Fatalf("Failed to run V3 migration: %v", err)
|
||||
}
|
||||
v3Duration := time.Since(startTime)
|
||||
logger.Infof("[Database] V3 迁移检查完成, 耗时: %v", v3Duration)
|
||||
|
||||
// 执行表结构同步
|
||||
migrateStart := time.Now()
|
||||
if err := database.Migrate(); err != nil {
|
||||
logger.Fatalf("Failed to migrate database: %v", err)
|
||||
}
|
||||
migrateDuration := time.Since(migrateStart)
|
||||
logger.Infof("[Database] 表结构同步完成, 耗时: %v", migrateDuration)
|
||||
logger.Infof("[Database] 数据库总初始化耗时: %v", time.Since(startTime))
|
||||
}
|
||||
|
||||
func (a *App) initRouter() {
|
||||
ctrls := router.RegisterControllers()
|
||||
a.Router = router.Setup(ctrls)
|
||||
}
|
||||
|
||||
func (a *App) Run() {
|
||||
addr := fmt.Sprintf("%s:%d", a.Config.Server.Host, a.Config.Server.Port)
|
||||
logger.Infof("Starting server on %s", addr)
|
||||
a.Router.Run(addr)
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package bootstrap
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
)
|
||||
|
||||
// SendInternalRequest 向常驻后台主服务安全发送内部通信请求
|
||||
// relPath 传入相对内部接口路径 (如: "/internal/tasks/execute/xxx"),方法内部会自动补充完整的协议、端口及 "/api/v1" 前缀,
|
||||
// 并自动获取 security.secret 密钥种入 X-Internal-Token 头部。
|
||||
func SendInternalRequest(method, relPath string, payload interface{}) ([]byte, int, error) {
|
||||
appCfg := services.GetConfig()
|
||||
if appCfg == nil {
|
||||
return nil, 0, fmt.Errorf("加载系统配置失败")
|
||||
}
|
||||
|
||||
relPath = strings.TrimPrefix(relPath, "/")
|
||||
url := fmt.Sprintf("http://127.0.0.1:%d/api/v1/%s", appCfg.Server.Port, relPath)
|
||||
|
||||
var bodyReader io.Reader
|
||||
if payload != nil {
|
||||
jsonData, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("序列化请求负载失败: %v", err)
|
||||
}
|
||||
bodyReader = bytes.NewBuffer(jsonData)
|
||||
}
|
||||
|
||||
settings := services.NewSettingsService()
|
||||
secret := settings.Get("security", "secret")
|
||||
|
||||
req, err := http.NewRequest(method, url, bodyReader)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("创建 HTTP 请求失败: %v", err)
|
||||
}
|
||||
|
||||
if payload != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
req.Header.Set("X-Internal-Token", secret)
|
||||
|
||||
client := &http.Client{Timeout: 5 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("网络连接失败,请确保任务池常驻后台服务正在运行中: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
bodyBytes, err := io.ReadAll(resp.Body)
|
||||
return bodyBytes, resp.StatusCode, err
|
||||
}
|
||||
Vendored
+94
@@ -0,0 +1,94 @@
|
||||
package cache
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
// siteCache 站点设置内存缓存
|
||||
var (
|
||||
siteCache = make(map[string]string)
|
||||
siteCacheMu sync.RWMutex
|
||||
siteCacheInit bool
|
||||
)
|
||||
|
||||
// LoadSiteCache 从数据库加载站点设置到缓存
|
||||
func LoadSiteCache() {
|
||||
siteCacheMu.Lock()
|
||||
defer siteCacheMu.Unlock()
|
||||
|
||||
// 先填充默认值
|
||||
if defaults, ok := constant.DefaultSettings[constant.SectionSite]; ok {
|
||||
for k, v := range defaults {
|
||||
siteCache[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
// 从数据库加载覆盖
|
||||
var settings []models.Setting
|
||||
database.DB.Where("section = ?", constant.SectionSite).Find(&settings)
|
||||
for _, setting := range settings {
|
||||
siteCache[setting.Key] = string(setting.Value)
|
||||
}
|
||||
siteCacheInit = true
|
||||
}
|
||||
|
||||
// ensureSiteCache 确保站点缓存已初始化
|
||||
func ensureSiteCache() {
|
||||
siteCacheMu.RLock()
|
||||
init := siteCacheInit
|
||||
siteCacheMu.RUnlock()
|
||||
|
||||
if !init {
|
||||
LoadSiteCache()
|
||||
}
|
||||
}
|
||||
|
||||
// GetSiteCache 从缓存获取站点设置
|
||||
func GetSiteCache(key string) string {
|
||||
ensureSiteCache()
|
||||
|
||||
siteCacheMu.RLock()
|
||||
defer siteCacheMu.RUnlock()
|
||||
|
||||
if val, ok := siteCache[key]; ok {
|
||||
return val
|
||||
}
|
||||
if def, ok := constant.DefaultSettings[constant.SectionSite][key]; ok {
|
||||
return def
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// SetSiteCache 更新缓存中的站点设置
|
||||
func SetSiteCache(key, value string) {
|
||||
siteCacheMu.Lock()
|
||||
siteCache[key] = value
|
||||
siteCacheMu.Unlock()
|
||||
}
|
||||
|
||||
// GetSiteCacheAll 获取整个站点设置缓存
|
||||
func GetSiteCacheAll() map[string]string {
|
||||
ensureSiteCache()
|
||||
|
||||
siteCacheMu.RLock()
|
||||
defer siteCacheMu.RUnlock()
|
||||
|
||||
result := make(map[string]string)
|
||||
for k, v := range siteCache {
|
||||
result[k] = v
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// SetSiteCacheBatch 批量更新缓存
|
||||
func SetSiteCacheBatch(values map[string]string) {
|
||||
siteCacheMu.Lock()
|
||||
for k, v := range values {
|
||||
siteCache[k] = v
|
||||
}
|
||||
siteCacheMu.Unlock()
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package constant
|
||||
|
||||
// CommandInfo 定义了终端可用命令的说明信息
|
||||
type CommandInfo struct {
|
||||
Name string
|
||||
Description string
|
||||
}
|
||||
|
||||
// Commands 是系统的可用业务命令说明列表
|
||||
var Commands = []CommandInfo{
|
||||
// {
|
||||
// Name: "server",
|
||||
// Description: "启动后台服务进程",
|
||||
// },
|
||||
{
|
||||
Name: "reposync",
|
||||
Description: "同步远程 Git 仓库或文件到本地",
|
||||
},
|
||||
{
|
||||
Name: "resetpwd",
|
||||
Description: "重置 admin 用户密码(需要二次确认)",
|
||||
},
|
||||
{
|
||||
Name: "restore",
|
||||
Description: "从本地 zip 文件中全量恢复系统级备份数据",
|
||||
},
|
||||
{
|
||||
Name: "builtininstall",
|
||||
Description: "为所有 mise 管理的 Node.js 和 Python 环境安装内建助手库",
|
||||
},
|
||||
{
|
||||
Name: "task",
|
||||
Description: "系统级任务的列表查询、触发运行、启停控制及状态查看",
|
||||
},
|
||||
{
|
||||
Name: "depinstall",
|
||||
Description: "一键补全指定任务执行日志中的缺失依赖包",
|
||||
},
|
||||
{
|
||||
Name: "version",
|
||||
Description: "查看当前系统版本号 (同 -v, -V)",
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,241 @@
|
||||
package constant
|
||||
|
||||
import "time"
|
||||
|
||||
const (
|
||||
|
||||
// DefaultRole 默认用户角色
|
||||
DefaultRole = "user"
|
||||
|
||||
// AdminRole 管理员角色
|
||||
AdminRole = "admin"
|
||||
|
||||
// DefaultTaskTimeout 默认任务超时时间(分钟)
|
||||
DefaultTaskTimeout = 30
|
||||
|
||||
// Settings Section 常量
|
||||
SectionSite = "site"
|
||||
SectionSystem = "system"
|
||||
SectionScheduler = "scheduler"
|
||||
SectionSecurity = "security"
|
||||
SectionNotify = "notify"
|
||||
|
||||
// Site Settings Key 常量
|
||||
KeyTitle = "title"
|
||||
KeySubtitle = "subtitle"
|
||||
KeyIcon = "icon"
|
||||
KeyPageSize = "page_size"
|
||||
KeyCookieDays = "cookie_days"
|
||||
KeyOpenapiToken = "openapi_token"
|
||||
KeyActiveWebUI = "active_webui"
|
||||
|
||||
// Security Settings Key 常量
|
||||
KeySecret = "secret"
|
||||
|
||||
// System Settings Key 常量
|
||||
KeyInitialized = "initialized"
|
||||
// KeyLogRetention = "log_retention" // Deprecated
|
||||
|
||||
// Log Retention Keys
|
||||
KeySystemNoticeDays = "system_notice_days"
|
||||
KeySystemNoticeMaxCount = "system_notice_max_count"
|
||||
KeyPushLogDays = "push_log_days"
|
||||
KeyPushLogMaxCount = "push_log_max_count"
|
||||
KeyLoginLogDays = "login_log_days"
|
||||
KeyLoginLogMaxCount = "login_log_max_count"
|
||||
KeySchedulerLogDays = "scheduler_log_days"
|
||||
KeySchedulerLogMaxCount = "scheduler_log_max_count"
|
||||
|
||||
// Scheduler Settings Key 常量
|
||||
KeyWorkerCount = "worker_count"
|
||||
KeyQueueSize = "queue_size"
|
||||
KeyRateInterval = "rate_interval"
|
||||
|
||||
// Notify Settings Key 常量
|
||||
KeyNotifyChannels = "channels"
|
||||
KeyNotifyEvents = "events"
|
||||
KeyNotifyToken = "notify_token"
|
||||
KeyNotifyPrefix = "notify_prefix"
|
||||
|
||||
// Notify Templates Keys
|
||||
KeyNotifyTemplateUserLoginTitle = "notify_template_user_login_title"
|
||||
KeyNotifyTemplateUserLoginText = "notify_template_user_login_text"
|
||||
KeyNotifyTemplateBruteForceLoginTitle = "notify_template_brute_force_login_title"
|
||||
KeyNotifyTemplateBruteForceLoginText = "notify_template_brute_force_login_text"
|
||||
KeyNotifyTemplatePasswordChangedTitle = "notify_template_password_changed_title"
|
||||
KeyNotifyTemplatePasswordChangedText = "notify_template_password_changed_text"
|
||||
KeyNotifyTemplateTaskSuccessTitle = "notify_template_task_success_title"
|
||||
KeyNotifyTemplateTaskSuccessText = "notify_template_task_success_text"
|
||||
KeyNotifyTemplateTaskFailedTitle = "notify_template_task_failed_title"
|
||||
KeyNotifyTemplateTaskFailedText = "notify_template_task_failed_text"
|
||||
KeyNotifyTemplateTaskTimeoutTitle = "notify_template_task_timeout_title"
|
||||
KeyNotifyTemplateTaskTimeoutText = "notify_template_task_timeout_text"
|
||||
|
||||
// 事件绑定类型
|
||||
BindingTypeSystem = "system"
|
||||
BindingTypeTask = "task"
|
||||
|
||||
// 系统事件类型
|
||||
EventUserLogin = "user_login"
|
||||
EventBruteForceLogin = "brute_force_login"
|
||||
EventPasswordChanged = "password_changed"
|
||||
|
||||
// 任务事件类型
|
||||
EventTaskSuccess = "task_success"
|
||||
EventTaskFailed = "task_failed"
|
||||
EventTaskTimeout = "task_timeout"
|
||||
EventTaskRunning = "task_running"
|
||||
EventTaskQueued = "task_queued"
|
||||
EventTaskCancelled = "task_cancelled"
|
||||
|
||||
// 其他事件类型
|
||||
EventSystemNotice = "system_notice"
|
||||
EventSchedulerLog = "scheduler_log"
|
||||
EventNotifySent = "notify_sent"
|
||||
EventAppLogAdded = "app_log_added"
|
||||
|
||||
// WebSocket 消息类型
|
||||
WSTypeHeartbeat = "heartbeat"
|
||||
WSTypeHeartbeatAck = "heartbeat_ack"
|
||||
WSTypeTasks = "tasks"
|
||||
WSTypeTaskResult = "task_result"
|
||||
WSTypeTaskLog = "task_log"
|
||||
WSTypeExecute = "execute"
|
||||
WSTypeUpdate = "update"
|
||||
WSTypeDisconnect = "disconnect"
|
||||
WSTypeConnected = "connected"
|
||||
WSTypeDisabled = "disabled"
|
||||
WSTypeEnabled = "enabled"
|
||||
WSTypeFetchTasks = "fetch_tasks"
|
||||
WSTypeTaskHeartbeat = "task_heartbeat"
|
||||
WSTypeStop = "stop"
|
||||
|
||||
// 任务状态
|
||||
TaskStatusSuccess = "success"
|
||||
TaskStatusFailed = "failed"
|
||||
TaskStatusRunning = "running"
|
||||
TaskStatusPending = "pending"
|
||||
TaskStatusTimeout = "timeout"
|
||||
TaskStatusCancelled = "cancelled"
|
||||
TaskStatusQueued = "queued"
|
||||
|
||||
// 任务类型
|
||||
TaskTypeNormal = "task"
|
||||
TaskTypeRepo = "repo"
|
||||
|
||||
// 任务置顶类型
|
||||
PinTypeNone = "none"
|
||||
PinTypeTop = "top"
|
||||
|
||||
// 触发类型(新值 taskpool_startup;读取时兼容旧值 baihu_startup)
|
||||
TriggerTypeCron = "cron"
|
||||
TriggerTypeTaskPoolStartup = "taskpool_startup"
|
||||
|
||||
// Agent 状态
|
||||
AgentStatusOnline = "online"
|
||||
AgentStatusOffline = "offline"
|
||||
|
||||
// AppLog 分类
|
||||
LogCategoryDefault = "default"
|
||||
LogCategorySystemNotice = "system_notice"
|
||||
LogCategoryPushLog = "push_log"
|
||||
LogCategoryLoginLog = "login_log"
|
||||
LogCategorySchedulerLog = "scheduler_log"
|
||||
|
||||
// AppLog 级别
|
||||
LogLevelInfo = "info"
|
||||
LogLevelWarning = "warning"
|
||||
LogLevelError = "error"
|
||||
|
||||
// AppLog 状态
|
||||
LogStatusUnread = "unread"
|
||||
LogStatusRead = "read"
|
||||
LogStatusSuccess = "success"
|
||||
LogStatusFailed = "failed"
|
||||
|
||||
// Env Type
|
||||
EnvTypeNormal = "normal"
|
||||
EnvTypeSecret = "secret"
|
||||
|
||||
// Relation Types
|
||||
RelationTypeTaskTag = "task_tag"
|
||||
RelationTypeTaskEnv = "task_env"
|
||||
RelationTypeEnvTag = "env_tag"
|
||||
|
||||
// WebSocket 安全常量
|
||||
// PongWait 收到 pong 的超时时间
|
||||
PongWait = 60 * time.Second
|
||||
// PingPeriod 发送 ping 的周期
|
||||
PingPeriod = (PongWait * 9) / 10
|
||||
// MaxMessageSize 允许的最大消息大小
|
||||
MaxMessageSize = 1024 * 1024 // 1MB
|
||||
// MaxLogSize 允许的最大日志大小 (保留末尾 10MB)
|
||||
MaxLogSize = 10 * 1024 * 1024 // 10MB
|
||||
|
||||
// ScriptsDirPlaceholder 脚本目录占位符
|
||||
ScriptsDirPlaceholder = "$SCRIPTS_DIR$"
|
||||
)
|
||||
|
||||
// CookieName Cookie 名称
|
||||
var CookieName = "BHToken"
|
||||
|
||||
// TablePrefix 表前缀,从配置文件读取
|
||||
var TablePrefix string
|
||||
|
||||
// Runtime 数据库配置快照,用于需要单独启动内部子进程(如 reposync)时显式透传数据库连接信息,
|
||||
// 避免主进程启动阶段清理环境变量后,子进程意外回退到默认 sqlite 配置。
|
||||
var (
|
||||
RuntimeDBType string
|
||||
RuntimeDBHost string
|
||||
RuntimeDBPort int
|
||||
RuntimeDBUser string
|
||||
RuntimeDBPassword string
|
||||
RuntimeDBName string
|
||||
RuntimeDBPath string
|
||||
RuntimeDBDSN string
|
||||
RuntimeDBTablePrefix string
|
||||
RuntimeDBSSLMode string
|
||||
)
|
||||
|
||||
// Secret JWT和密码salt密钥,运行中自动从数据库加载
|
||||
var Secret string
|
||||
|
||||
// DemoMode 演示模式,从环境变量读取
|
||||
var DemoMode bool
|
||||
|
||||
// DefaultIcon 默认站点图标
|
||||
var DefaultIcon = `<svg t="1766107903919" class="icon" viewBox="0 0 1024 1024" version="1.1" xmlns="http://www.w3.org/2000/svg" p-id="1942" width="200" height="200"><path d="M884.992 273.05984c4.10624 0 5.0688-2.36544 2.10944-5.25312 0 0-64.28672-65.55648-111.7696-75.55072-47.48288-9.984-45.47584-59.37152-80.56832-75.02848-72.0896-32.16384-158.6176-34.3552-158.6176-34.3552s-91.37152-4.92544-138.752-9.89184c-19.46624-2.03776-54.46656-9.58464-54.46656-9.58464-4.0448-0.84992-10.63936-1.19808-14.66368-0.44032 0 0-30.63808-0.07168-44.30848 43.35616-8.8576 28.11904 1.792 104.79616 1.792 104.79616 1.46432 12.27776-3.42016 30.21824-10.72128 40.20224 0 0-36.77184 46.03904-58.9312 100.34176-22.15936 54.30272 118.15936 145.05984 208.98816 205.27104C507.82208 613.89824 502.03648 743.424 502.03648 743.424s-74.19904-96.75776-194.00704-156.9792C188.2112 526.22336 150.30272 442.23488 150.30272 442.23488c-2.89792-5.45792-5.94944-4.94592-6.71744 1.21856 0 0-15.1552 91.61728 16.25088 147.8144 70.92224 126.88384 141.74208 112.88576 197.03808 183.27552s54.272 164.9152 54.272 164.9152-97.62816-141.29152-235.66336-205.55776c-91.53536-61.27616-74.7008-125.91104-85.49376-101.85728-10.79296 24.05376 26.73664 192.41984 65.40288 222.2592 80.57856 62.18752 94.16704 101.98016 94.16704 101.98016h175.53408S544.512 814.85824 572.928 725.73952c44.41088-139.30496 40.20224-191.26272 40.20224-191.26272 0.08192-8.22272 6.81984-14.19264 14.98112-13.29152 0 0 46.45888 4.64896 66.64192 9.68704 23.53152 5.86752 55.35744 26.20416 55.35744 26.20416 3.49184 2.14016 7.39328 0.68608 8.69376-3.1744l34.4576-102.77888c1.30048-3.8912-0.63488-5.51936-4.352-3.75808 0 0-45.37344 25.09824-88.17664 12.1856-20.15232-6.08256-59.60704-14.82752-74.69056-32.6656-16.95744-20.03968-15.59552-71.3728 26.66496-79.21664 48.0256-8.9088 33.13664 15.14496 65.91488 24.64768 27.42272 7.95648 22.29248-1.69984 26.69568 5.21216 1.67936 2.63168 0.38912 32.65536 0.38912 32.65536-0.21504 6.144 3.39968 7.84384 8.0384 3.79904l41.89184-36.46464c17.37728-11.24352 30.86336-0.57344 48.24064-11.81696 14.73536-9.53344 29.58336-43.66336 29.58336-43.66336 4.48512-9.24672 0.60416-20.41856-8.63232-24.92416l-42.58816-20.80768c-3.69664-1.80224-3.38944-3.26656 0.74752-3.26656h62.0032zM422.54336 123.87328s-49.88928 50.7392-74.5472 50.66752c-24.65792-0.07168-33.24928-56.12544-3.92192-56.12544 18.00192 0 76.32896 0.21504 76.32896 0.21504 4.13696 0.02048 5.0688 2.36544 2.14016 5.24288z m123.09504 249.64096s-3.31776-25.53856-33.16736-40.05888c-29.8496-14.52032-52.5312-41.13408-59.648-68.95616-12.1856-54.8864 48.29184-104.192 48.29184-104.192s-30.1056 73.6768 3.5328 106.60864c54.272 53.12512 40.99072 106.5984 40.99072 106.5984z m155.56608-164.46464c-7.94624 10.32192-9.60512 37.92896-59.2384 20.736-49.63328-17.2032-71.00416-54.5792-71.00416-54.5792-2.27328-3.40992-0.79872-6.49216 3.328-6.8096 0 0 43.55072-5.03808 80.06656 6.44096 24.91392 7.82336 54.79424 23.88992 46.848 34.21184z" fill="#272636" p-id="1943"></path><path d="M366.30528 259.26656c1.05472-1.76128 1.51552-1.57696 1.19808 0.43008 0 0-7.55712 19.0464 22.8864 87.63392 27.0848 61.02016 87.49056 68.7104 118.66112 98.23232 55.808 52.82816 53.52448 123.45344 53.52448 123.45344s-25.82528-49.85856-85.98528-78.83776-92.6208-50.52416-127.7952-85.77024c-45.02528-45.12768 17.5104-145.14176 17.5104-145.14176zM500.0704 961.44384h134.49216s95.8464-138.07616 46.68416-291.25632c-16.19968-50.46272-43.45856-80.00512-43.45856-80.00512-5.18144-6.38976-8.89856-4.88448-8.448 3.34848 0 0 9.15456 99.80928-19.44576 194.00704-28.60032 94.208-109.824 173.90592-109.824 173.90592zM681.61536 956.8768h105.24672s22.38464-73.5232 16.61952-130.21184c-9.30816-91.57632-48.31232-132.72064-48.31232-132.72064-3.80928-4.77184-6.0928-3.67616-5.2224 2.38592 0 0 16.75264 89.1904-14.4896 150.1184-30.9248 60.30336-53.84192 110.42816-53.84192 110.42816zM869.30432 811.35616c-2.88768-2.9184-4.80256-1.95584-4.38272 2.10944 0 0 6.79936 44.99456-7.76192 81.22368-10.12736 25.1904-28.91776 58.60352-28.91776 58.60352h107.17184s5.34528-50.31936-19.89632-83.97824c-11.56096-23.99232-46.21312-57.9584-46.21312-57.9584z" fill="#272636" p-id="1944"></path></svg>`
|
||||
|
||||
// DefaultSettings 默认系统设置
|
||||
var DefaultSettings = map[string]map[string]string{
|
||||
SectionSite: {
|
||||
KeyTitle: "任务池",
|
||||
KeySubtitle: "极致轻量、高性能的自动化任务调度平台",
|
||||
KeyIcon: DefaultIcon,
|
||||
KeyPageSize: "10",
|
||||
KeyCookieDays: "7",
|
||||
KeyActiveWebUI: "default",
|
||||
},
|
||||
SectionScheduler: {
|
||||
KeyWorkerCount: "4",
|
||||
KeyQueueSize: "100",
|
||||
KeyRateInterval: "200",
|
||||
},
|
||||
SectionNotify: {
|
||||
KeyNotifyPrefix: "[任务池]",
|
||||
// Login
|
||||
KeyNotifyTemplateUserLoginTitle: "用户登录(成功/失败)",
|
||||
KeyNotifyTemplateUserLoginText: "用户 {{username}} 在 IP {{ip}} 登录{{status_label}}\n{{message}}",
|
||||
KeyNotifyTemplateBruteForceLoginTitle: "系统安全警告",
|
||||
KeyNotifyTemplateBruteForceLoginText: "检测到 IP {{ip}} 正在尝试暴力破解用户 {{username}}",
|
||||
KeyNotifyTemplatePasswordChangedTitle: "账户安全通知",
|
||||
KeyNotifyTemplatePasswordChangedText: "用户 {{username}} 刚刚修改了密码",
|
||||
// Task
|
||||
KeyNotifyTemplateTaskSuccessTitle: "任务[{{task_name}}] 成功",
|
||||
KeyNotifyTemplateTaskSuccessText: "任务 #{{task_id}} {{task_name}}\n状态: 成功\n耗时: {{duration}}ms\n执行结果: {{output}}",
|
||||
KeyNotifyTemplateTaskFailedTitle: "任务[{{task_name}}] 失败",
|
||||
KeyNotifyTemplateTaskFailedText: "任务 #{{task_id}} {{task_name}}\n状态: 失败\n执行时间: {{start_time}}\n原因: {{error}}\n最后输出: {{output}}",
|
||||
KeyNotifyTemplateTaskTimeoutTitle: "任务[{{task_name}}] 超时",
|
||||
KeyNotifyTemplateTaskTimeoutText: "任务 #{{task_id}} {{task_name}}\n状态: 超时\n耗时: {{duration}}ms\n最后输出: {{output}}",
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
package constant
|
||||
|
||||
const (
|
||||
// CookieActiveInterconnectNodeID 穿越状态下标识目标子节点 ID 的 Cookie 键名
|
||||
CookieActiveInterconnectNodeID = "active_interconnect_node_id"
|
||||
|
||||
// SectionInterconnect 互联设置分组
|
||||
SectionInterconnect = "interconnect"
|
||||
|
||||
// 互联设置相关 Key
|
||||
KeyInterconnectToken = "interconnect_token"
|
||||
KeyInterconnectParentURL = "interconnect_parent_url"
|
||||
KeyInterconnectParentToken = "interconnect_parent_token"
|
||||
KeyInterconnectRole = "interconnect_role"
|
||||
|
||||
// 互联角色
|
||||
InterconnectRoleMaster = "master"
|
||||
InterconnectRoleChild = "child"
|
||||
|
||||
// 互联系统事件
|
||||
EventInterconnectChildStatus = "interconnect_child_status"
|
||||
)
|
||||
@@ -0,0 +1,21 @@
|
||||
package constant
|
||||
|
||||
const (
|
||||
// IDSchemaSignature 数据库表结构指纹的固定 ID
|
||||
IDSchemaSignature = "sys_schema_sig"
|
||||
|
||||
// KeySchemaSignature 数据库表结构指纹,用于判断是否需要全量自动建表
|
||||
KeySchemaSignature = "schema_signature"
|
||||
|
||||
// KeyTaskEnvsMigrated 任务环境变量迁移标记 (v2版本强制重跑过)
|
||||
KeyTaskEnvsMigrated = "task_envs_migrated_v2"
|
||||
|
||||
// KeyTaskTagsMigrated 任务标签迁移标记 (v2版本强制重跑过)
|
||||
KeyTaskTagsMigrated = "task_tags_migrated_v2"
|
||||
|
||||
// 以下是旧表或字段名,用于前置结构迁移时重命名或删除
|
||||
TableMigrateQlTokens = "ql_tokens"
|
||||
ColumnMigrateQlTokenCode = "code"
|
||||
ColumnMigrateQlTokenToken = "token"
|
||||
ColumnMigrateDependencyType = "type"
|
||||
)
|
||||
@@ -0,0 +1,9 @@
|
||||
package constant
|
||||
|
||||
// MainstreamMisePlugins 主流的 mise 插件列表
|
||||
var MainstreamMisePlugins = []string{
|
||||
"python", "node", "go", "rust", "ruby", "java", "php",
|
||||
"deno", "bun", "zig", "dotnet", "elixir", "erlang",
|
||||
"crystal", "lua", "julia", "nim", "perl", "scala",
|
||||
"kotlin", "clojure", "dart", "flutter", "terraform",
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
package constant
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
)
|
||||
|
||||
var (
|
||||
// ConfigPath 配置文件路径
|
||||
ConfigPath string
|
||||
|
||||
// DataDir 数据目录
|
||||
DataDir string
|
||||
|
||||
// DefaultDBPath 默认数据库路径
|
||||
DefaultDBPath string
|
||||
|
||||
// WebDistDir 前端构建目录
|
||||
WebDistDir string
|
||||
|
||||
// ScriptsWorkDir 脚本工作目录
|
||||
ScriptsWorkDir string
|
||||
)
|
||||
|
||||
func init() {
|
||||
rootDir := ResolveAppRootDir()
|
||||
ConfigPath = filepath.Clean(filepath.Join(rootDir, "configs", "config.ini"))
|
||||
DataDir = filepath.Clean(filepath.Join(rootDir, "data"))
|
||||
DefaultDBPath = filepath.Clean(filepath.Join(rootDir, "data", "taskpool.db"))
|
||||
WebDistDir = filepath.Clean(filepath.Join(rootDir, "web", "dist"))
|
||||
ScriptsWorkDir = filepath.Clean(filepath.Join(rootDir, "data", "scripts"))
|
||||
}
|
||||
|
||||
// ResolveAppRootDir 获取应用程序的绝对根目录路径。
|
||||
func ResolveAppRootDir() string {
|
||||
// 1. 检查当前工作目录(CWD)及其上级目录
|
||||
if cwd, err := os.Getwd(); err == nil {
|
||||
dir := cwd
|
||||
for {
|
||||
if _, err := os.Stat(filepath.Join(dir, "configs", "config.ini")); err == nil {
|
||||
return dir
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, "go.mod")); err == nil {
|
||||
return dir
|
||||
}
|
||||
parent := filepath.Dir(dir)
|
||||
if parent == dir {
|
||||
break
|
||||
}
|
||||
dir = parent
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 检查当前可执行文件路径及其上级目录
|
||||
if exe, err := os.Executable(); err == nil {
|
||||
dir := filepath.Dir(exe)
|
||||
for {
|
||||
if _, err := os.Stat(filepath.Join(dir, "configs", "config.ini")); err == nil {
|
||||
return dir
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dir, "go.mod")); err == nil {
|
||||
return dir
|
||||
}
|
||||
parent := filepath.Dir(dir)
|
||||
if parent == dir {
|
||||
break
|
||||
}
|
||||
dir = parent
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 兜底回退到当前工作目录
|
||||
if cwd, err := os.Getwd(); err == nil {
|
||||
return cwd
|
||||
}
|
||||
return "."
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package constant
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
_ "embed"
|
||||
"encoding/json"
|
||||
"math/rand"
|
||||
"sync"
|
||||
)
|
||||
|
||||
//go:embed sentence1-10000.json
|
||||
var sentenceData []byte
|
||||
|
||||
type Sentence struct {
|
||||
Name string `json:"name"`
|
||||
From string `json:"from"`
|
||||
}
|
||||
|
||||
var (
|
||||
lineCount int
|
||||
lineOffsets []int64
|
||||
sentenceOnce sync.Once
|
||||
)
|
||||
|
||||
// initSentences 初始化:统计行数和记录每行偏移
|
||||
func initSentences() {
|
||||
sentenceOnce.Do(func() {
|
||||
data := sentenceData
|
||||
var offset int64 = 0
|
||||
for {
|
||||
// 记录当前行的起始位置
|
||||
lineOffsets = append(lineOffsets, offset)
|
||||
lineCount++
|
||||
|
||||
// 寻找下一个换行符
|
||||
idx := bytes.IndexByte(data[offset:], '\n')
|
||||
if idx == -1 {
|
||||
// 最后一行没有换行符
|
||||
break
|
||||
}
|
||||
|
||||
// 移动到下一行的起始位置
|
||||
offset += int64(idx) + 1
|
||||
if offset >= int64(len(data)) {
|
||||
break
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// GetRandomSentence 随机获取一条古诗词
|
||||
func GetRandomSentence() string {
|
||||
initSentences()
|
||||
if lineCount <= 0 {
|
||||
return "欢迎使用任务池"
|
||||
}
|
||||
|
||||
targetIndex := rand.Intn(lineCount)
|
||||
start := lineOffsets[targetIndex]
|
||||
|
||||
// 确定当前行的结束位置
|
||||
var end int64
|
||||
if targetIndex < lineCount-1 {
|
||||
end = lineOffsets[targetIndex+1]
|
||||
} else {
|
||||
end = int64(len(sentenceData))
|
||||
}
|
||||
|
||||
// 提取行并清理两端的空白字符(包括 \r, \n)
|
||||
line := bytes.TrimSpace(sentenceData[start:end])
|
||||
if len(line) == 0 {
|
||||
return "欢迎使用任务池"
|
||||
}
|
||||
|
||||
var sData []string
|
||||
if err := json.Unmarshal(line, &sData); err == nil && len(sData) >= 1 {
|
||||
name := sData[0]
|
||||
from := ""
|
||||
if len(sData) >= 2 {
|
||||
from = sData[1]
|
||||
}
|
||||
|
||||
if from != "" {
|
||||
return "\"" + name + "\"—— " + from
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
return "欢迎使用任务池"
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,12 @@
|
||||
package constant
|
||||
|
||||
import "time"
|
||||
|
||||
// 构建时注入的变量
|
||||
var (
|
||||
Version = "dev"
|
||||
BuildTime = "unknown"
|
||||
)
|
||||
|
||||
// 程序启动时间
|
||||
var StartTime = time.Now()
|
||||
@@ -0,0 +1,746 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
var agentUpgrader = websocket.Upgrader{
|
||||
CheckOrigin: utils.CheckWSOrigin,
|
||||
}
|
||||
|
||||
// AgentController Agent 控制器
|
||||
type AgentController struct {
|
||||
agentService *services.AgentService
|
||||
wsManager *services.AgentWSManager
|
||||
settingsService *services.SettingsService
|
||||
}
|
||||
|
||||
// NewAgentController 创建 Agent 控制器
|
||||
func NewAgentController(settingsService *services.SettingsService) *AgentController {
|
||||
return &AgentController{
|
||||
agentService: services.NewAgentService(),
|
||||
wsManager: services.GetAgentWSManager(),
|
||||
settingsService: settingsService,
|
||||
}
|
||||
}
|
||||
|
||||
// List 获取 Agent 列表
|
||||
func (c *AgentController) List(ctx *gin.Context) {
|
||||
agents := c.agentService.List()
|
||||
utils.Success(ctx, vo.ToAgentVOListFromModels(agents))
|
||||
}
|
||||
|
||||
// getActiveSchedulerConfig 获取 Agent 的实际调度配置(若为空或零值,则使用系统默认的 settings)
|
||||
func (c *AgentController) getActiveSchedulerConfig(agent *models.Agent) map[string]interface{} {
|
||||
workerCount := agent.SchedulerConfig.WorkerCount
|
||||
queueSize := agent.SchedulerConfig.QueueSize
|
||||
rateInterval := int(agent.SchedulerConfig.RateInterval / time.Millisecond)
|
||||
strictQueue := agent.SchedulerConfig.StrictQueue
|
||||
|
||||
// 如果未配置(WorkerCount <= 0),则使用全局系统设置
|
||||
if workerCount <= 0 {
|
||||
workerCount = getIntSetting(c.settingsService, constant.SectionScheduler, constant.KeyWorkerCount, 4)
|
||||
queueSize = getIntSetting(c.settingsService, constant.SectionScheduler, constant.KeyQueueSize, 100)
|
||||
rateInterval = getIntSetting(c.settingsService, constant.SectionScheduler, constant.KeyRateInterval, 200)
|
||||
strictQueue = false
|
||||
}
|
||||
|
||||
return map[string]interface{}{
|
||||
"worker_count": workerCount,
|
||||
"queue_size": queueSize,
|
||||
"rate_interval": rateInterval,
|
||||
"strict_queue": strictQueue,
|
||||
}
|
||||
}
|
||||
|
||||
// Update 更新 Agent
|
||||
func (c *AgentController) Update(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Description string `json:"description"`
|
||||
Enabled bool `json:"enabled"`
|
||||
SchedulerConfig *vo.AgentSchedulerConfigVO `json:"scheduler_config"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取旧状态
|
||||
oldAgent := c.agentService.GetByID(id)
|
||||
if oldAgent == nil {
|
||||
utils.NotFound(ctx, "Agent 不存在")
|
||||
return
|
||||
}
|
||||
wasEnabled := utils.DerefBool(oldAgent.Enabled, true)
|
||||
|
||||
var schedulerConfig models.AgentSchedulerConfig
|
||||
if req.SchedulerConfig != nil {
|
||||
schedulerConfig.WorkerCount = req.SchedulerConfig.WorkerCount
|
||||
schedulerConfig.QueueSize = req.SchedulerConfig.QueueSize
|
||||
schedulerConfig.RateInterval = time.Duration(req.SchedulerConfig.RateInterval) * time.Millisecond
|
||||
schedulerConfig.Verbose = req.SchedulerConfig.Verbose
|
||||
schedulerConfig.StrictQueue = req.SchedulerConfig.StrictQueue
|
||||
}
|
||||
|
||||
if err := c.agentService.Update(id, req.Name, req.Description, req.Enabled, schedulerConfig); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 如果启用状态发生变化,通知 Agent
|
||||
if wasEnabled != req.Enabled {
|
||||
if req.Enabled {
|
||||
// 启用:发送任务列表
|
||||
c.wsManager.SendToAgent(id, services.WSTypeEnabled, map[string]interface{}{
|
||||
"message": "Agent 已启用",
|
||||
})
|
||||
// 发送任务列表
|
||||
c.wsManager.BroadcastTasks(id)
|
||||
} else {
|
||||
// 禁用:发送禁用消息,Agent 收到后清空任务
|
||||
c.wsManager.SendToAgent(id, services.WSTypeDisabled, map[string]interface{}{
|
||||
"message": "Agent 已禁用",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 推送最新的调度配置给 Agent (如果 Agent 在线)
|
||||
if req.Enabled {
|
||||
// 重新加载已更新的 Agent 信息以获取正确的 SchedulerConfig
|
||||
updatedAgent := c.agentService.GetByID(id)
|
||||
if updatedAgent != nil {
|
||||
c.wsManager.SendToAgent(id, services.WSTypeConnected, map[string]interface{}{
|
||||
"agent_id": id,
|
||||
"name": req.Name,
|
||||
"scheduler_config": c.getActiveSchedulerConfig(updatedAgent),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "更新成功")
|
||||
}
|
||||
|
||||
// Delete 删除 Agent
|
||||
func (c *AgentController) Delete(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.agentService.Delete(id); err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "删除成功")
|
||||
}
|
||||
|
||||
// RegenerateToken 重新生成 Token
|
||||
func (c *AgentController) RegenerateToken(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
token, err := c.agentService.RegenerateToken(id)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"token": token})
|
||||
}
|
||||
|
||||
// ========== Agent API(供 Agent 调用)==========
|
||||
|
||||
// Register Agent 注册(无需认证)
|
||||
func (c *AgentController) Register(ctx *gin.Context) {
|
||||
var req models.AgentRegisterRequest
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" {
|
||||
utils.BadRequest(ctx, "名称不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
ip := ctx.ClientIP()
|
||||
agent, token, err := c.agentService.Register(&req, ip)
|
||||
if err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{
|
||||
"agent_id": agent.ID,
|
||||
"token": token,
|
||||
"message": "注册成功",
|
||||
})
|
||||
}
|
||||
|
||||
// Heartbeat Agent 心跳
|
||||
func (c *AgentController) Heartbeat(ctx *gin.Context) {
|
||||
token := c.getAgentToken(ctx)
|
||||
if token == "" {
|
||||
utils.Unauthorized(ctx, "缺少认证 Token")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Version string `json:"version"`
|
||||
BuildTime string `json:"build_time"`
|
||||
Hostname string `json:"hostname"`
|
||||
OS string `json:"os"`
|
||||
Arch string `json:"arch"`
|
||||
AutoUpdate bool `json:"auto_update"`
|
||||
}
|
||||
ctx.ShouldBindJSON(&req)
|
||||
|
||||
ip := ctx.ClientIP()
|
||||
agent, err := c.agentService.Heartbeat(token, ip, req.Version, req.BuildTime, req.Hostname, req.OS, req.Arch)
|
||||
if err != nil {
|
||||
utils.Unauthorized(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 检查是否需要更新
|
||||
latestVersion := c.agentService.GetLatestVersion()
|
||||
needUpdate := c.agentService.CheckNeedUpdate(req.Version, req.BuildTime)
|
||||
forceUpdate := agent.ForceUpdate
|
||||
|
||||
// 如果强制更新已触发,重置标志
|
||||
if forceUpdate && needUpdate {
|
||||
c.agentService.ClearForceUpdate(agent.ID)
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{
|
||||
"agent_id": agent.ID,
|
||||
"name": agent.Name,
|
||||
"need_update": needUpdate,
|
||||
"force_update": forceUpdate,
|
||||
"latest_version": latestVersion,
|
||||
})
|
||||
}
|
||||
|
||||
// GetTasks Agent 获取任务列表
|
||||
func (c *AgentController) GetTasks(ctx *gin.Context) {
|
||||
token := c.getAgentToken(ctx)
|
||||
if token == "" {
|
||||
utils.Unauthorized(ctx, "缺少认证 Token")
|
||||
return
|
||||
}
|
||||
|
||||
// 先尝试通过 token 查找 Agent
|
||||
agent := c.agentService.GetByToken(token)
|
||||
|
||||
// 如果找不到,尝试验证令牌并通过 machine_id 查找
|
||||
if agent == nil {
|
||||
machineID := ctx.GetHeader("X-Machine-ID")
|
||||
if machineID != "" {
|
||||
// 验证令牌是否有效
|
||||
if _, err := c.agentService.ValidateToken(token); err == nil {
|
||||
// 令牌有效,尝试通过 machine_id 查找 Agent
|
||||
agent = c.agentService.GetByMachineID(machineID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if agent == nil {
|
||||
utils.Unauthorized(ctx, "无效的 Token")
|
||||
return
|
||||
}
|
||||
|
||||
if !utils.DerefBool(agent.Enabled, true) {
|
||||
utils.Forbidden(ctx, "Agent 已禁用")
|
||||
return
|
||||
}
|
||||
|
||||
tasks := c.agentService.GetTasks(agent.ID)
|
||||
utils.Success(ctx, gin.H{
|
||||
"agent_id": agent.ID,
|
||||
"tasks": tasks,
|
||||
})
|
||||
}
|
||||
|
||||
// ReportResult Agent 上报执行结果
|
||||
func (c *AgentController) ReportResult(ctx *gin.Context) {
|
||||
token := c.getAgentToken(ctx)
|
||||
if token == "" {
|
||||
utils.Unauthorized(ctx, "缺少认证 Token")
|
||||
return
|
||||
}
|
||||
|
||||
agent := c.agentService.GetByToken(token)
|
||||
if agent == nil {
|
||||
utils.Unauthorized(ctx, "无效的 Token")
|
||||
return
|
||||
}
|
||||
|
||||
if !utils.DerefBool(agent.Enabled, true) {
|
||||
utils.Forbidden(ctx, "Agent 已禁用")
|
||||
return
|
||||
}
|
||||
|
||||
var result models.AgentTaskResult
|
||||
if err := ctx.ShouldBindJSON(&result); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
result.AgentID = agent.ID
|
||||
|
||||
if err := c.agentService.ReportResult(&result); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "上报成功")
|
||||
}
|
||||
|
||||
// getAgentToken 从请求头获取 Agent Token
|
||||
func (c *AgentController) getAgentToken(ctx *gin.Context) string {
|
||||
auth := ctx.GetHeader("Authorization")
|
||||
if auth == "" {
|
||||
return ""
|
||||
}
|
||||
// Bearer <token>
|
||||
parts := strings.SplitN(auth, " ", 2)
|
||||
if len(parts) == 2 && parts[0] == "Bearer" {
|
||||
return parts[1]
|
||||
}
|
||||
return auth
|
||||
}
|
||||
|
||||
// Download 下载 Agent 程序
|
||||
func (c *AgentController) Download(ctx *gin.Context) {
|
||||
osType := ctx.DefaultQuery("os", "linux")
|
||||
arch := ctx.DefaultQuery("arch", "amd64")
|
||||
|
||||
data, filename, err := c.agentService.GetAgentBinary(osType, arch)
|
||||
if err != nil {
|
||||
utils.NotFound(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
ctx.Header("Content-Disposition", "attachment; filename="+filename)
|
||||
ctx.Header("Content-Type", "application/gzip")
|
||||
ctx.Header("Content-Length", strconv.Itoa(len(data)))
|
||||
ctx.Data(200, "application/gzip", data)
|
||||
}
|
||||
|
||||
// GetVersion 获取 Agent 最新版本信息
|
||||
func (c *AgentController) GetVersion(ctx *gin.Context) {
|
||||
version := c.agentService.GetLatestVersion()
|
||||
platforms := c.agentService.GetAvailablePlatforms()
|
||||
|
||||
utils.Success(ctx, gin.H{
|
||||
"version": version,
|
||||
"platforms": platforms,
|
||||
})
|
||||
}
|
||||
|
||||
// ForceUpdate 强制更新指定 Agent
|
||||
func (c *AgentController) ForceUpdate(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.agentService.SetForceUpdate(id); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "已标记强制更新,Agent 下次心跳时将自动更新")
|
||||
}
|
||||
|
||||
// ========== WebSocket ==========
|
||||
|
||||
// WSConnect Agent WebSocket 连接
|
||||
func (c *AgentController) WSConnect(ctx *gin.Context) {
|
||||
// 添加 panic 恢复
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logger.Errorf("[AgentWS] WSConnect panic: %v", r)
|
||||
ctx.JSON(http.StatusInternalServerError, gin.H{"error": "服务器内部错误"})
|
||||
}
|
||||
}()
|
||||
|
||||
ip := ctx.ClientIP()
|
||||
|
||||
// 打印请求信息用于调试
|
||||
logger.Infof("[AgentWS] 收到连接请求: IP=%s, URL=%s", ip, ctx.Request.URL.String())
|
||||
|
||||
// 检查 IP 限流
|
||||
if allowed, reason := c.wsManager.CheckRateLimit(ip); !allowed {
|
||||
logger.Warnf("[AgentWS] IP %s 被限流: %s", ip, reason)
|
||||
ctx.JSON(http.StatusTooManyRequests, gin.H{"error": reason})
|
||||
return
|
||||
}
|
||||
|
||||
token := ctx.Query("token")
|
||||
if token == "" {
|
||||
c.wsManager.RecordConnectFail(ip)
|
||||
logger.Warnf("[AgentWS] 连接失败: 缺少 token, IP=%s", ip)
|
||||
ctx.JSON(http.StatusUnauthorized, gin.H{"error": "缺少 token"})
|
||||
return
|
||||
}
|
||||
|
||||
machineID := ctx.Query("machine_id")
|
||||
logger.Infof("[AgentWS] Token: %s..., MachineID: %s...", token[:8], machineID[:16])
|
||||
|
||||
isNewAgent := false
|
||||
|
||||
// 先尝试用 token 查找已有 Agent
|
||||
agent := c.agentService.GetByToken(token)
|
||||
logger.Infof("[AgentWS] GetByToken 结果: agent=%v", agent != nil)
|
||||
|
||||
// 如果没找到,尝试用令牌注册(会检查 machine_id 是否已存在)
|
||||
if agent == nil {
|
||||
logger.Infof("[AgentWS] 尝试注册新 Agent")
|
||||
var err error
|
||||
agent, isNewAgent, err = c.agentService.RegisterByToken(token, machineID, ip)
|
||||
if err != nil {
|
||||
c.wsManager.RecordConnectFail(ip)
|
||||
logger.Warnf("[AgentWS] 注册失败: %v, IP=%s, token=%s", err, ip, token[:8]+"...")
|
||||
ctx.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
logger.Infof("[AgentWS] 注册成功: Agent #%s, isNew=%v", agent.ID, isNewAgent)
|
||||
}
|
||||
|
||||
if !utils.DerefBool(agent.Enabled, true) {
|
||||
c.wsManager.RecordConnectFail(ip)
|
||||
logger.Warnf("[AgentWS] Agent #%s 已禁用, IP=%s", agent.ID, ip)
|
||||
ctx.JSON(http.StatusForbidden, gin.H{"error": "Agent 已禁用"})
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof("[AgentWS] 准备升级连接: Agent #%s, IP=%s", agent.ID, ip)
|
||||
conn, err := agentUpgrader.Upgrade(ctx.Writer, ctx.Request, nil)
|
||||
if err != nil {
|
||||
logger.Errorf("[AgentWS] 升级连接失败: %v, Agent #%s, IP=%s", err, agent.ID, ip)
|
||||
return
|
||||
}
|
||||
conn.SetReadLimit(constant.MaxMessageSize)
|
||||
|
||||
// 连接成功,重置失败计数
|
||||
c.wsManager.RecordConnectSuccess(ip)
|
||||
|
||||
// 注册连接
|
||||
ac := c.wsManager.Register(agent.ID, conn, ip)
|
||||
|
||||
// 更新 Agent 状态
|
||||
c.agentService.Heartbeat(token, ip, "", "", "", "", "")
|
||||
|
||||
// 获取调度配置并发送连接成功消息(包含注册状态和调度配置)
|
||||
schedCfg := c.getActiveSchedulerConfig(agent)
|
||||
c.wsManager.SendToAgent(agent.ID, services.WSTypeConnected, map[string]interface{}{
|
||||
"agent_id": agent.ID,
|
||||
"name": agent.Name,
|
||||
"is_new_agent": isNewAgent,
|
||||
"machine_id": machineID,
|
||||
"scheduler_config": schedCfg,
|
||||
})
|
||||
|
||||
logger.Infof("[AgentWS] Agent #%s 连接成功 (配置: %v)", agent.ID, schedCfg)
|
||||
|
||||
// 启动读写协程
|
||||
go c.wsWritePump(ac)
|
||||
go c.wsReadPump(ac, agent)
|
||||
|
||||
// 主动推送任务列表
|
||||
go c.wsManager.BroadcastTasks(agent.ID)
|
||||
}
|
||||
|
||||
// wsReadPump 读取消息
|
||||
func (c *AgentController) wsReadPump(ac *services.AgentConnection, agent *models.Agent) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logger.Errorf("[AgentWS] Agent #%s wsReadPump panic: %v", agent.ID, r)
|
||||
}
|
||||
logger.Infof("[AgentWS] Agent #%s wsReadPump 退出", agent.ID)
|
||||
c.wsManager.Unregister(agent.ID, ac)
|
||||
}()
|
||||
|
||||
// 检查连接是否有效(可能是旧连接被新连接替换)
|
||||
if ac == nil || ac.IsClosed() {
|
||||
return
|
||||
}
|
||||
|
||||
ac.SetReadDeadline(time.Now().Add(90 * time.Second))
|
||||
// 注意:SetPongHandler 需要直接访问 Conn,但这里我们在连接建立后立即设置
|
||||
// 所以是安全的,因为此时连接还没有被其他 goroutine 关闭
|
||||
ac.Conn.SetPongHandler(func(string) error {
|
||||
ac.SetReadDeadline(time.Now().Add(90 * time.Second))
|
||||
return nil
|
||||
})
|
||||
|
||||
for {
|
||||
_, message, err := ac.ReadMessage()
|
||||
if err != nil {
|
||||
logger.Warnf("[AgentWS] Agent #%s 读取错误: %v", agent.ID, err)
|
||||
break
|
||||
}
|
||||
|
||||
var msg services.WSMessage
|
||||
if err := json.Unmarshal(message, &msg); err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
c.handleWSMessage(ac, agent, &msg)
|
||||
}
|
||||
}
|
||||
|
||||
// wsWritePump 写入消息
|
||||
func (c *AgentController) wsWritePump(ac *services.AgentConnection) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logger.Errorf("[AgentWS] Agent #%s wsWritePump panic: %v", ac.AgentID, r)
|
||||
}
|
||||
logger.Infof("[AgentWS] Agent #%s wsWritePump 退出", ac.AgentID)
|
||||
}()
|
||||
|
||||
ticker := time.NewTicker(30 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case message, ok := <-ac.Send:
|
||||
if !ok {
|
||||
logger.Warnf("[AgentWS] Agent #%s Send channel 已关闭", ac.AgentID)
|
||||
return
|
||||
}
|
||||
if ac.IsClosed() {
|
||||
logger.Warnf("[AgentWS] Agent #%s 连接已关闭(write)", ac.AgentID)
|
||||
return
|
||||
}
|
||||
if err := ac.WriteMessage(message); err != nil {
|
||||
logger.Warnf("[AgentWS] Agent #%s 写入消息失败: %v", ac.AgentID, err)
|
||||
return
|
||||
}
|
||||
case <-ticker.C:
|
||||
if ac.IsClosed() {
|
||||
return
|
||||
}
|
||||
if err := ac.WritePing(); err != nil {
|
||||
logger.Warnf("[AgentWS] Agent #%s 发送 Ping 失败: %v", ac.AgentID, err)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// handleWSMessage 处理 WebSocket 消息
|
||||
func (c *AgentController) handleWSMessage(ac *services.AgentConnection, agent *models.Agent, msg *services.WSMessage) {
|
||||
switch msg.Type {
|
||||
case services.WSTypeHeartbeat:
|
||||
c.handleHeartbeat(ac, agent, msg.Data)
|
||||
|
||||
case services.WSTypeTaskResult:
|
||||
c.handleTaskResult(agent, msg.Data)
|
||||
|
||||
case services.WSTypeTaskLog:
|
||||
c.handleTaskLog(agent, msg.Data)
|
||||
|
||||
case services.WSTypeFetchTasks:
|
||||
c.handleFetchTasks(agent)
|
||||
|
||||
case services.WSTypeTaskHeartbeat: // 任务心跳
|
||||
c.handleTaskHeartbeat(agent, msg.Data)
|
||||
}
|
||||
}
|
||||
|
||||
// handleTaskHeartbeat 处理任务心跳
|
||||
func (c *AgentController) handleTaskHeartbeat(_ *models.Agent, data json.RawMessage) {
|
||||
var req struct {
|
||||
LogID string `json:"log_id"`
|
||||
Duration int64 `json:"duration"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &req); err != nil {
|
||||
logger.Errorf("[AgentWS] 解析心跳消息失败: %v", err)
|
||||
return
|
||||
}
|
||||
if req.LogID != "" {
|
||||
logger.Infof("[AgentWS] 收到任务心跳: LogID=%s, Duration=%dms", req.LogID, req.Duration)
|
||||
c.agentService.UpdateTaskDuration(req.LogID, req.Duration)
|
||||
}
|
||||
}
|
||||
|
||||
// handleFetchTasks 处理 Agent 请求任务列表
|
||||
func (c *AgentController) handleFetchTasks(agent *models.Agent) {
|
||||
tasks := c.agentService.GetTasks(agent.ID)
|
||||
c.wsManager.SendToAgent(agent.ID, services.WSTypeTasks, map[string]interface{}{
|
||||
"tasks": tasks,
|
||||
})
|
||||
logger.Infof("[AgentWS] Agent #%s 请求任务列表,返回 %d 个任务", agent.ID, len(tasks))
|
||||
}
|
||||
|
||||
// handleHeartbeat 处理心跳
|
||||
func (c *AgentController) handleHeartbeat(ac *services.AgentConnection, agent *models.Agent, data json.RawMessage) {
|
||||
var req struct {
|
||||
Version string `json:"version"`
|
||||
BuildTime string `json:"build_time"`
|
||||
Hostname string `json:"hostname"`
|
||||
OS string `json:"os"`
|
||||
Arch string `json:"arch"`
|
||||
AutoUpdate bool `json:"auto_update"`
|
||||
}
|
||||
json.Unmarshal(data, &req)
|
||||
|
||||
ac.UpdatePing()
|
||||
|
||||
// 更新 Agent 信息(使用连接时保存的 IP)
|
||||
c.agentService.Heartbeat(agent.Token, ac.IP, req.Version, req.BuildTime, req.Hostname, req.OS, req.Arch)
|
||||
|
||||
// 检查是否需要更新
|
||||
latestVersion := c.agentService.GetLatestVersion()
|
||||
needUpdate := c.agentService.CheckNeedUpdate(req.Version, req.BuildTime)
|
||||
forceUpdate := agent.ForceUpdate
|
||||
|
||||
if forceUpdate && needUpdate {
|
||||
c.agentService.ClearForceUpdate(agent.ID)
|
||||
}
|
||||
|
||||
// 发送心跳响应
|
||||
response := map[string]interface{}{
|
||||
"agent_id": agent.ID,
|
||||
"name": agent.Name,
|
||||
"need_update": needUpdate,
|
||||
"force_update": forceUpdate,
|
||||
"latest_version": latestVersion,
|
||||
}
|
||||
c.wsManager.SendToAgent(agent.ID, services.WSTypeHeartbeatAck, response)
|
||||
}
|
||||
|
||||
// handleTaskResult 处理任务结果
|
||||
func (c *AgentController) handleTaskResult(agent *models.Agent, data json.RawMessage) {
|
||||
var result models.AgentTaskResult
|
||||
if err := json.Unmarshal(data, &result); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
result.AgentID = agent.ID
|
||||
c.agentService.ReportResult(&result)
|
||||
}
|
||||
|
||||
// handleTaskLog 处理 Agent 发送的实时日志
|
||||
func (c *AgentController) handleTaskLog(_ *models.Agent, data json.RawMessage) {
|
||||
var logMsg struct {
|
||||
LogID string `json:"log_id"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &logMsg); err != nil {
|
||||
logger.Errorf("[AgentWS] 解析日志消息失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
tl := tasks.GetActiveLog(logMsg.LogID)
|
||||
if tl != nil {
|
||||
tl.Write([]byte(logMsg.Content))
|
||||
} else {
|
||||
logger.Warnf("[AgentWS] 收到任务日志 but could not find active TinyLog: LogID=%s, ContentSize=%d", logMsg.LogID, len(logMsg.Content))
|
||||
}
|
||||
}
|
||||
|
||||
// NotifyTaskUpdate 通知 Agent 任务更新
|
||||
func (c *AgentController) NotifyTaskUpdate(agentID string) {
|
||||
c.wsManager.BroadcastTasks(agentID)
|
||||
}
|
||||
|
||||
// ========== 令牌管理 ==========
|
||||
|
||||
// ListTokens 获取令牌列表
|
||||
func (c *AgentController) ListTokens(ctx *gin.Context) {
|
||||
tokens := c.agentService.ListTokens()
|
||||
utils.Success(ctx, vo.ToAgentTokenVOListFromModels(tokens))
|
||||
}
|
||||
|
||||
// CreateToken 创建令牌
|
||||
func (c *AgentController) CreateToken(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Remark string `json:"remark"`
|
||||
MaxUses int `json:"max_uses"`
|
||||
ExpiresAt string `json:"expires_at"` // 格式: 2006-01-02 15:04:05
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
var expiresAt *time.Time
|
||||
if req.ExpiresAt != "" {
|
||||
t, err := time.ParseInLocation("2006-01-02 15:04:05", req.ExpiresAt, time.Local)
|
||||
if err != nil {
|
||||
utils.BadRequest(ctx, "过期时间格式错误")
|
||||
return
|
||||
}
|
||||
expiresAt = &t
|
||||
}
|
||||
|
||||
token, err := c.agentService.CreateToken(req.Remark, req.MaxUses, expiresAt)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, vo.ToAgentTokenVO(token))
|
||||
}
|
||||
|
||||
// DeleteToken 删除令牌
|
||||
func (c *AgentController) DeleteToken(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.agentService.DeleteToken(id); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "删除成功")
|
||||
}
|
||||
|
||||
// getIntSetting 辅助方法
|
||||
func getIntSetting(s *services.SettingsService, section, key string, defaultVal int) int {
|
||||
val := s.Get(section, key)
|
||||
|
||||
if val == "" {
|
||||
return defaultVal
|
||||
}
|
||||
if result, err := strconv.Atoi(val); err == nil {
|
||||
return result
|
||||
}
|
||||
return defaultVal
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type AppLogController struct {
|
||||
appLogService *services.AppLogService
|
||||
}
|
||||
|
||||
func NewAppLogController() *AppLogController {
|
||||
return &AppLogController{
|
||||
appLogService: services.NewAppLogService(),
|
||||
}
|
||||
}
|
||||
|
||||
// GetLogs 获取应用日志列表
|
||||
func (ac *AppLogController) GetLogs(c *gin.Context) {
|
||||
p := utils.ParsePagination(c)
|
||||
category := c.Query("category")
|
||||
status := c.Query("status")
|
||||
level := c.Query("level")
|
||||
keyword := c.Query("keyword")
|
||||
|
||||
logs, total, err := ac.appLogService.List(category, status, level, p.Page, p.PageSize, keyword)
|
||||
if err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
utils.PaginatedResponse(c, logs, total, p)
|
||||
}
|
||||
|
||||
// MarkAsRead 标记已读
|
||||
func (ac *AppLogController) MarkAsRead(c *gin.Context) {
|
||||
var req struct {
|
||||
ID string `json:"id"`
|
||||
Category string `json:"category"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.ID != "" {
|
||||
if err := ac.appLogService.MarkAsRead(req.ID); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
} else if req.Category != "" {
|
||||
if err := ac.appLogService.MarkAllAsRead(req.Category); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
} else {
|
||||
utils.BadRequest(c, "id 或 category 必须提供")
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(c, "标记成功")
|
||||
}
|
||||
|
||||
// ClearLogs 清理日志
|
||||
func (ac *AppLogController) ClearLogs(c *gin.Context) {
|
||||
var req struct {
|
||||
Category string `json:"category"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if err := ac.appLogService.Clear(req.Category); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(c, "清理成功")
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/eventbus"
|
||||
"github.com/engigu/taskpool/internal/middleware"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type AuthController struct {
|
||||
userService *services.UserService
|
||||
settingsService *services.SettingsService
|
||||
loginLogService *services.LoginLogService
|
||||
}
|
||||
|
||||
type loginAttempt struct {
|
||||
Count int
|
||||
LastAttempt time.Time
|
||||
}
|
||||
|
||||
var loginAttempts sync.Map
|
||||
|
||||
func init() {
|
||||
// 定期清理过期的登录尝试统计,防止内存溢出
|
||||
go func() {
|
||||
ticker := time.NewTicker(30 * time.Minute)
|
||||
for range ticker.C {
|
||||
loginAttempts.Range(func(key, value any) bool {
|
||||
attempt := value.(*loginAttempt)
|
||||
if time.Since(attempt.LastAttempt) > 10*time.Minute {
|
||||
loginAttempts.Delete(key)
|
||||
}
|
||||
return true
|
||||
})
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func NewAuthController(userService *services.UserService, settingsService *services.SettingsService, loginLogService *services.LoginLogService) *AuthController {
|
||||
return &AuthController{
|
||||
userService: userService,
|
||||
settingsService: settingsService,
|
||||
loginLogService: loginLogService,
|
||||
}
|
||||
}
|
||||
|
||||
func (ac *AuthController) Login(c *gin.Context) {
|
||||
var req struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
|
||||
ip := c.ClientIP()
|
||||
userAgent := c.GetHeader("User-Agent")
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 暴力破解防御
|
||||
if val, ok := loginAttempts.Load(ip); ok {
|
||||
attempt := val.(*loginAttempt)
|
||||
if attempt.Count >= 5 && time.Since(attempt.LastAttempt) < time.Minute {
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventBruteForceLogin,
|
||||
Payload: map[string]interface{}{
|
||||
"ip": ip,
|
||||
"username": req.Username,
|
||||
"userAgent": userAgent,
|
||||
},
|
||||
})
|
||||
utils.TooManyRequests(c, "尝试次数过多,请一分钟后再试")
|
||||
return
|
||||
}
|
||||
// 如果距离上次尝试已超过一分钟,重置计数
|
||||
if time.Since(attempt.LastAttempt) >= time.Minute {
|
||||
loginAttempts.Delete(ip)
|
||||
}
|
||||
}
|
||||
|
||||
user := ac.userService.GetUserByUsername(req.Username)
|
||||
if user == nil || !ac.userService.ValidatePassword(user, req.Password) {
|
||||
// 记录失败尝试
|
||||
val, _ := loginAttempts.LoadOrStore(ip, &loginAttempt{Count: 0, LastAttempt: time.Now()})
|
||||
attempt := val.(*loginAttempt)
|
||||
attempt.Count++
|
||||
attempt.LastAttempt = time.Now()
|
||||
|
||||
// 记录登录失败日志
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventUserLogin,
|
||||
Payload: map[string]interface{}{
|
||||
"ip": ip,
|
||||
"username": req.Username,
|
||||
"userAgent": userAgent,
|
||||
"status": "failed",
|
||||
"message": "用户名或密码错误",
|
||||
},
|
||||
})
|
||||
utils.Unauthorized(c, "用户名或密码错误")
|
||||
return
|
||||
}
|
||||
|
||||
// 登录成功,清除尝试记录
|
||||
loginAttempts.Delete(ip)
|
||||
|
||||
// 获取 cookie 过期天数
|
||||
expireDays := 7
|
||||
if days := ac.settingsService.Get(constant.SectionSite, constant.KeyCookieDays); days != "" {
|
||||
if d, err := strconv.Atoi(days); err == nil && d > 0 {
|
||||
expireDays = d
|
||||
}
|
||||
}
|
||||
|
||||
// 生成 token
|
||||
token, err := utils.GenerateToken(user.ID, user.Username, user.TokenVersion, expireDays, constant.Secret)
|
||||
if err != nil {
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventUserLogin,
|
||||
Payload: map[string]interface{}{
|
||||
"ip": ip,
|
||||
"username": req.Username,
|
||||
"userAgent": userAgent,
|
||||
"status": "failed",
|
||||
"message": "Token生成失败",
|
||||
},
|
||||
})
|
||||
utils.ServerError(c, "登录失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 设置 Cookie
|
||||
middleware.SetAuthCookie(c, token, expireDays)
|
||||
|
||||
// 记录登录成功日志
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventUserLogin,
|
||||
Payload: map[string]interface{}{
|
||||
"ip": ip,
|
||||
"username": req.Username,
|
||||
"userAgent": userAgent,
|
||||
"status": "success",
|
||||
"message": "登录成功",
|
||||
},
|
||||
})
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"user": user.Username,
|
||||
})
|
||||
}
|
||||
|
||||
func (ac *AuthController) Logout(c *gin.Context) {
|
||||
if userID, exists := c.Get("userID"); exists {
|
||||
ac.userService.InvalidateUserTokens(userID.(string))
|
||||
}
|
||||
middleware.ClearAuthCookie(c)
|
||||
utils.SuccessMsg(c, "退出成功")
|
||||
}
|
||||
|
||||
func (ac *AuthController) GetCurrentUser(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
user, err := ac.userService.GetUserByID(userID)
|
||||
if err != nil {
|
||||
utils.Unauthorized(c, "会话无效")
|
||||
return
|
||||
}
|
||||
utils.Success(c, gin.H{
|
||||
"username": user.Username,
|
||||
"role": user.Role,
|
||||
})
|
||||
}
|
||||
|
||||
func (ac *AuthController) Register(c *gin.Context) {
|
||||
/*
|
||||
var req struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Email string `json:"email" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 安全性:强制设定角色为 user,防止注册时篡改角色为 admin
|
||||
user := ac.userService.CreateUser(req.Username, req.Password, req.Email, constant.DefaultRole)
|
||||
utils.Success(c, vo.ToUserVO(user))
|
||||
*/
|
||||
utils.BadRequest(c, "注册功能已关闭")
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type DashboardController struct {
|
||||
executorService *tasks.ExecutorService
|
||||
}
|
||||
|
||||
func NewDashboardController(executorService *tasks.ExecutorService) *DashboardController {
|
||||
return &DashboardController{
|
||||
executorService: executorService,
|
||||
}
|
||||
}
|
||||
|
||||
type StatsResponse struct {
|
||||
Tasks int64 `json:"tasks"`
|
||||
TodayExecs int64 `json:"today_execs"`
|
||||
Envs int64 `json:"envs"`
|
||||
Logs int64 `json:"logs"`
|
||||
Scheduled int `json:"scheduled"`
|
||||
Running int `json:"running"`
|
||||
}
|
||||
|
||||
func (dc *DashboardController) GetStats(c *gin.Context) {
|
||||
var taskCount, envCount, logCount, todayExecs int64
|
||||
|
||||
database.DB.Model(&models.Task{}).Count(&taskCount)
|
||||
database.DB.Model(&models.EnvironmentVariable{}).Count(&envCount)
|
||||
database.DB.Model(&models.TaskLog{}).Count(&logCount)
|
||||
|
||||
// 今日执行总数
|
||||
today := time.Now().Format("2006-01-02")
|
||||
database.DB.Model(&models.SendStats{}).Where("day = ?", today).Select("COALESCE(SUM(num), 0)").Scan(&todayExecs)
|
||||
|
||||
// 调度统计:本地调度 + Agent 调度
|
||||
// 本地调度:agent_id 为 NULL 且 enabled = true 的任务
|
||||
localScheduled := dc.executorService.GetScheduledCount()
|
||||
|
||||
// Agent 调度:agent_id 不为 NULL 且 enabled = true 的任务
|
||||
var agentScheduled int64
|
||||
database.DB.Model(&models.Task{}).
|
||||
Where("agent_id IS NOT NULL AND enabled = ?", true).
|
||||
Count(&agentScheduled)
|
||||
|
||||
totalScheduled := localScheduled + int(agentScheduled)
|
||||
|
||||
// 正在运行:目前只能统计本地运行的任务
|
||||
// Agent 端的运行状态需要通过心跳上报(未来优化)
|
||||
running := dc.executorService.GetRunningCount()
|
||||
|
||||
stats := StatsResponse{
|
||||
Tasks: taskCount,
|
||||
TodayExecs: todayExecs,
|
||||
Envs: envCount,
|
||||
Logs: logCount,
|
||||
Scheduled: totalScheduled,
|
||||
Running: running,
|
||||
}
|
||||
|
||||
utils.Success(c, stats)
|
||||
}
|
||||
|
||||
// GetSentence 获取随机古诗词
|
||||
func (dc *DashboardController) GetSentence(c *gin.Context) {
|
||||
utils.Success(c, gin.H{
|
||||
"sentence": constant.GetRandomSentence(),
|
||||
})
|
||||
}
|
||||
|
||||
// DailyStats 每日统计数据
|
||||
type DailyStats struct {
|
||||
Day string `json:"day"`
|
||||
Total int `json:"total"`
|
||||
Success int `json:"success"`
|
||||
Failed int `json:"failed"`
|
||||
}
|
||||
|
||||
// GetSendStats 获取发送统计
|
||||
func (dc *DashboardController) GetSendStats(c *gin.Context) {
|
||||
// 获取天数参数,默认30天
|
||||
days := 30
|
||||
if d := c.Query("days"); d != "" {
|
||||
if parsed, err := utils.ParseInt(d); err == nil && parsed > 0 && parsed <= 90 {
|
||||
days = parsed
|
||||
}
|
||||
}
|
||||
|
||||
// 获取日期范围
|
||||
now := time.Now()
|
||||
startDay := now.AddDate(0, 0, -(days - 1)).Format("2006-01-02")
|
||||
|
||||
var stats []models.SendStats
|
||||
database.DB.Where("day >= ?", startDay).Find(&stats)
|
||||
|
||||
// 按日期聚合
|
||||
dayMap := make(map[string]*DailyStats)
|
||||
for _, s := range stats {
|
||||
if _, ok := dayMap[s.Day]; !ok {
|
||||
dayMap[s.Day] = &DailyStats{Day: s.Day}
|
||||
}
|
||||
ds := dayMap[s.Day]
|
||||
ds.Total += s.Num
|
||||
if s.Status == constant.TaskStatusSuccess {
|
||||
ds.Success += s.Num
|
||||
} else {
|
||||
ds.Failed += s.Num
|
||||
}
|
||||
}
|
||||
|
||||
// 填充缺失的日期
|
||||
result := make([]DailyStats, 0, days)
|
||||
for i := days - 1; i >= 0; i-- {
|
||||
day := now.AddDate(0, 0, -i).Format("2006-01-02")
|
||||
if ds, ok := dayMap[day]; ok {
|
||||
result = append(result, *ds)
|
||||
} else {
|
||||
result = append(result, DailyStats{Day: day})
|
||||
}
|
||||
}
|
||||
|
||||
// 按日期排序
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
return result[i].Day < result[j].Day
|
||||
})
|
||||
|
||||
utils.Success(c, result)
|
||||
}
|
||||
|
||||
// TaskStats 任务执行统计
|
||||
type TaskStats struct {
|
||||
TaskID string `json:"task_id"`
|
||||
TaskName string `json:"task_name"`
|
||||
Count int `json:"count"`
|
||||
}
|
||||
|
||||
// GetTaskStats 获取任务执行占比
|
||||
func (dc *DashboardController) GetTaskStats(c *gin.Context) {
|
||||
// 获取天数参数,默认30天
|
||||
days := 30
|
||||
if d := c.Query("days"); d != "" {
|
||||
if parsed, err := utils.ParseInt(d); err == nil && parsed > 0 && parsed <= 90 {
|
||||
days = parsed
|
||||
}
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
startDay := now.AddDate(0, 0, -(days - 1)).Format("2006-01-02")
|
||||
|
||||
// 按 task_id 聚合统计
|
||||
var results []struct {
|
||||
TaskID string
|
||||
Total int
|
||||
}
|
||||
database.DB.Model(&models.SendStats{}).
|
||||
Select("task_id, SUM(num) as total").
|
||||
Where("day >= ?", startDay).
|
||||
Group("task_id").
|
||||
Order("total DESC").
|
||||
Find(&results)
|
||||
|
||||
// 获取任务名称
|
||||
taskIDs := make([]string, 0, len(results))
|
||||
for _, r := range results {
|
||||
taskIDs = append(taskIDs, r.TaskID)
|
||||
}
|
||||
|
||||
var tasks []models.Task
|
||||
if len(taskIDs) > 0 {
|
||||
database.DB.Where("id IN ?", taskIDs).Find(&tasks)
|
||||
}
|
||||
taskNameMap := make(map[string]string)
|
||||
for _, t := range tasks {
|
||||
taskNameMap[t.ID] = t.Name
|
||||
}
|
||||
|
||||
// 构建结果
|
||||
stats := make([]TaskStats, 0, len(results))
|
||||
for _, r := range results {
|
||||
name := taskNameMap[r.TaskID]
|
||||
if name == "" {
|
||||
name = "未知任务"
|
||||
}
|
||||
stats = append(stats, TaskStats{
|
||||
TaskID: r.TaskID,
|
||||
TaskName: name,
|
||||
Count: r.Total,
|
||||
})
|
||||
}
|
||||
|
||||
utils.Success(c, stats)
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type DataController struct {
|
||||
dataService *services.DataService
|
||||
taskController *TaskController
|
||||
envController *EnvController
|
||||
}
|
||||
|
||||
func NewDataController(tc *TaskController, ec *EnvController) *DataController {
|
||||
return &DataController{
|
||||
dataService: services.NewDataService(),
|
||||
taskController: tc,
|
||||
envController: ec,
|
||||
}
|
||||
}
|
||||
|
||||
// ExportBusinessData 导出业务数据
|
||||
func (dc *DataController) ExportBusinessData(c *gin.Context) {
|
||||
var req struct {
|
||||
TaskIDs []string `json:"task_ids"`
|
||||
EnvIDs []string `json:"env_ids"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
exportData := dc.dataService.ExportBusinessData(req.TaskIDs, req.EnvIDs)
|
||||
utils.Success(c, exportData)
|
||||
}
|
||||
|
||||
// ImportBusinessData 导入业务数据
|
||||
func (dc *DataController) ImportBusinessData(c *gin.Context) {
|
||||
var req models.ExportData
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.Version == "" {
|
||||
utils.BadRequest(c, "无效的导入数据格式")
|
||||
return
|
||||
}
|
||||
|
||||
// 停止相关的定时任务
|
||||
if len(req.Tasks) > 0 {
|
||||
for _, task := range req.Tasks {
|
||||
dc.taskController.executorService.RemoveCronTask(task.ID)
|
||||
dc.taskController.executorService.GetScheduler().StopTask(task.ID)
|
||||
}
|
||||
}
|
||||
|
||||
// 导入数据
|
||||
if err := dc.dataService.ImportBusinessData(&req); err != nil {
|
||||
utils.ServerError(c, "导入失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 重新启动任务和通知相关的代理
|
||||
if len(req.Tasks) > 0 {
|
||||
for i := range req.Tasks {
|
||||
task := &req.Tasks[i]
|
||||
if utils.DerefBool(task.Enabled, true) && (task.AgentID == nil || *task.AgentID == "") {
|
||||
dc.taskController.executorService.AddCronTask(task)
|
||||
}
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
dc.taskController.agentWSManager.BroadcastTasks(*task.AgentID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "导入成功")
|
||||
}
|
||||
@@ -0,0 +1,421 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/deps"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type DependencyController struct {
|
||||
service *services.DependencyService
|
||||
}
|
||||
|
||||
func NewDependencyController() *DependencyController {
|
||||
return &DependencyController{
|
||||
service: services.NewDependencyService(),
|
||||
}
|
||||
}
|
||||
|
||||
// List 获取依赖列表
|
||||
func (c *DependencyController) List(ctx *gin.Context) {
|
||||
language := ctx.Query("language")
|
||||
langVersion := ctx.Query("lang_version")
|
||||
deps, err := c.service.List(language, langVersion)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取依赖列表失败")
|
||||
return
|
||||
}
|
||||
vos := vo.ToDependencyVOListFromModels(deps)
|
||||
utils.Success(ctx, vos)
|
||||
}
|
||||
|
||||
// Create 添加依赖
|
||||
func (c *DependencyController) Create(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Version string `json:"version"`
|
||||
Language string `json:"language" binding:"required"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
dep := &models.Dependency{
|
||||
Name: req.Name,
|
||||
Version: req.Version,
|
||||
Language: req.Language,
|
||||
LangVersion: req.LangVersion,
|
||||
Remark: req.Remark,
|
||||
}
|
||||
|
||||
if err := c.service.Create(dep); err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, vo.ToDependencyVO(dep))
|
||||
}
|
||||
|
||||
// Delete 删除依赖
|
||||
func (c *DependencyController) Delete(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.service.Delete(id); err != nil {
|
||||
utils.ServerError(ctx, "删除失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "删除成功")
|
||||
}
|
||||
|
||||
func (c *DependencyController) Install(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Version string `json:"version"`
|
||||
Language string `json:"language"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
language := req.Language
|
||||
if language == "" {
|
||||
language = ctx.Query("language")
|
||||
}
|
||||
langVersion := req.LangVersion
|
||||
if langVersion == "" {
|
||||
langVersion = ctx.Query("lang_version")
|
||||
}
|
||||
|
||||
dep := &models.Dependency{
|
||||
Name: req.Name,
|
||||
Version: req.Version,
|
||||
Language: language,
|
||||
LangVersion: langVersion,
|
||||
Remark: req.Remark,
|
||||
}
|
||||
|
||||
err := c.service.Install(dep)
|
||||
// 无论成功失败,都同步记录日志
|
||||
c.service.Create(dep)
|
||||
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(ctx, "安装成功")
|
||||
}
|
||||
|
||||
// GetInstallCommand 获取安装命令
|
||||
func (c *DependencyController) GetInstallCommand(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Version string `json:"version"`
|
||||
Language string `json:"language"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
language := req.Language
|
||||
if language == "" {
|
||||
language = ctx.Query("language")
|
||||
}
|
||||
langVersion := req.LangVersion
|
||||
if langVersion == "" {
|
||||
langVersion = ctx.Query("lang_version")
|
||||
}
|
||||
|
||||
dep := &models.Dependency{
|
||||
Name: req.Name,
|
||||
Version: req.Version,
|
||||
Language: language,
|
||||
LangVersion: langVersion,
|
||||
}
|
||||
|
||||
cmd, err := c.service.GetInstallCommand(dep)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"command": cmd})
|
||||
}
|
||||
|
||||
// GetReinstallAllCommand 获取全部重装命令
|
||||
func (c *DependencyController) GetReinstallAllCommand(ctx *gin.Context) {
|
||||
language := ctx.Query("language")
|
||||
langVersion := ctx.Query("lang_version")
|
||||
if language == "" {
|
||||
utils.BadRequest(ctx, "缺少 language 参数")
|
||||
return
|
||||
}
|
||||
|
||||
cmd, err := c.service.GetReinstallAllCommand(language, langVersion)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"command": cmd})
|
||||
}
|
||||
|
||||
// Uninstall 卸载依赖
|
||||
func (c *DependencyController) Uninstall(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
force := ctx.Query("force") == "true"
|
||||
|
||||
// 获取依赖信息
|
||||
deps, _ := c.service.List("", "")
|
||||
var dep *models.Dependency
|
||||
for i := range deps {
|
||||
if deps[i].ID == id {
|
||||
dep = &deps[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if dep == nil {
|
||||
utils.NotFound(ctx, "依赖不存在")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.service.Uninstall(dep); err != nil {
|
||||
if !force {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 卸载成功(或强制删除)后从数据库删除
|
||||
c.service.Delete(id)
|
||||
|
||||
utils.SuccessMsg(ctx, "卸载成功")
|
||||
}
|
||||
|
||||
// Reinstall 重新安装依赖
|
||||
func (c *DependencyController) Reinstall(ctx *gin.Context) {
|
||||
id := ctx.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(ctx, "无效的 ID")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取依赖信息
|
||||
deps, _ := c.service.List("", "")
|
||||
var dep *models.Dependency
|
||||
for i := range deps {
|
||||
if deps[i].ID == id {
|
||||
dep = &deps[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if dep == nil {
|
||||
utils.NotFound(ctx, "依赖不存在")
|
||||
return
|
||||
}
|
||||
|
||||
err := c.service.Install(dep)
|
||||
// 无论成功失败,都同步记录日志
|
||||
c.service.Create(dep)
|
||||
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(ctx, "重新安装成功")
|
||||
}
|
||||
|
||||
// ReinstallAll 重新安装所有依赖
|
||||
func (c *DependencyController) ReinstallAll(ctx *gin.Context) {
|
||||
language := ctx.Query("language")
|
||||
langVersion := ctx.Query("lang_version")
|
||||
if language == "" {
|
||||
utils.BadRequest(ctx, "缺少 language 参数")
|
||||
return
|
||||
}
|
||||
|
||||
deps, err := c.service.List(language, langVersion)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取依赖列表失败")
|
||||
return
|
||||
}
|
||||
|
||||
var failed []string
|
||||
for i := range deps {
|
||||
d := &deps[i]
|
||||
err := c.service.Install(d)
|
||||
if err != nil {
|
||||
failed = append(failed, d.Name)
|
||||
}
|
||||
// 无论成功失败,都同步记录日志到数据库
|
||||
c.service.Create(d)
|
||||
}
|
||||
|
||||
if len(failed) > 0 {
|
||||
utils.ServerError(ctx, "部分包安装失败: "+strings.Join(failed, ", "))
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(ctx, "全部重新安装成功")
|
||||
}
|
||||
|
||||
// GetInstalled 获取已安装的包
|
||||
func (c *DependencyController) GetInstalled(ctx *gin.Context) {
|
||||
language := ctx.Query("language")
|
||||
langVersion := ctx.Query("lang_version")
|
||||
if language == "" {
|
||||
utils.BadRequest(ctx, "缺少 language 参数")
|
||||
return
|
||||
}
|
||||
|
||||
packages, err := c.service.GetInstalledPackages(language, langVersion)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取已安装包失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, packages)
|
||||
}
|
||||
|
||||
// GetBatchInstallCommand 获取批量安装依赖包的命令
|
||||
func (c *DependencyController) GetBatchInstallCommand(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Items []struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Version string `json:"version"`
|
||||
Language string `json:"language" binding:"required"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
} `json:"items" binding:"required,gt=0"`
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: items 不能为空且必须包含 name 和 language")
|
||||
return
|
||||
}
|
||||
|
||||
var depsList []models.Dependency
|
||||
for _, item := range req.Items {
|
||||
depsList = append(depsList, models.Dependency{
|
||||
Name: item.Name,
|
||||
Version: item.Version,
|
||||
Language: item.Language,
|
||||
LangVersion: item.LangVersion,
|
||||
})
|
||||
}
|
||||
|
||||
cmd, err := c.service.GetBatchInstallCommand(depsList)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"command": cmd})
|
||||
}
|
||||
|
||||
// ParseAndImport 解析上传/粘贴的清单文件内容并批量导入至数据库
|
||||
func (c *DependencyController) ParseAndImport(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Language string `json:"language" binding:"required"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
ImportDB bool `json:"import_db"` // 是否持久化到数据库做可视化管理
|
||||
}
|
||||
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: language 和 content 必填")
|
||||
return
|
||||
}
|
||||
|
||||
// 1. 解析文本清单内容
|
||||
parsedDeps, err := deps.ParseManifest(req.Language, req.Content)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "清单文件解析失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if len(parsedDeps) == 0 {
|
||||
utils.BadRequest(ctx, "未解析到任何有效依赖包")
|
||||
return
|
||||
}
|
||||
|
||||
// 2. 补全语言和版本属性
|
||||
for i := range parsedDeps {
|
||||
parsedDeps[i].Language = req.Language
|
||||
parsedDeps[i].LangVersion = req.LangVersion
|
||||
}
|
||||
|
||||
// 3. 根据需求决定是否导入数据库
|
||||
var finalDeps []models.Dependency
|
||||
if req.ImportDB {
|
||||
imported, err := c.service.ImportDependencies(parsedDeps)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "导入依赖记录至数据库失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
finalDeps = imported
|
||||
} else {
|
||||
finalDeps = parsedDeps
|
||||
}
|
||||
|
||||
// 4. 为这一批包生成合并批量安装命令
|
||||
cmd, err := c.service.GetBatchInstallCommand(finalDeps)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "生成安装命令失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{
|
||||
"dependencies": vo.ToDependencyVOListFromModels(finalDeps),
|
||||
"command": cmd,
|
||||
})
|
||||
}
|
||||
|
||||
// GetDepInstallCommand 获取自动补全的命令,返回给前端执行
|
||||
func (c *DependencyController) GetDepInstallCommand(ctx *gin.Context) {
|
||||
logID := ctx.Query("log_id")
|
||||
if logID == "" {
|
||||
utils.BadRequest(ctx, "参数错误: log_id 不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
execPath, err := os.Executable()
|
||||
if err != nil {
|
||||
execPath = "taskpool" // 兜底
|
||||
}
|
||||
|
||||
// 构造命令,比如: "F:\workspace\taskpool\taskpool.exe" depinstall <log_id>
|
||||
cmdStr := fmt.Sprintf("%q depinstall %s", execPath, logID)
|
||||
utils.Success(ctx, gin.H{
|
||||
"command": cmdStr,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,389 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/relation"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type EnvController struct {
|
||||
envService *services.EnvService
|
||||
}
|
||||
|
||||
func NewEnvController(envService *services.EnvService) *EnvController {
|
||||
return &EnvController{envService: envService}
|
||||
}
|
||||
|
||||
// GetSecretStatus 获取加密秘钥状态
|
||||
// @Summary 获取加密秘钥状态
|
||||
// @Description 返回系统是否已配置加密秘钥
|
||||
// @Tags Env
|
||||
// @Produce json
|
||||
// @Success 200 {object} utils.Response{data=bool} "成功"
|
||||
// @Router /env/secret-status [get]
|
||||
// @Security BearerAuth
|
||||
func (ec *EnvController) GetSecretStatus(c *gin.Context) {
|
||||
utils.Success(c, utils.IsSecretKeySet())
|
||||
}
|
||||
|
||||
// CreateEnvVar 创建环境变量
|
||||
// @Summary 创建环境变量
|
||||
// @Description 创建一个新的环境变量
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param body body object true "环境变量信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.EnvVO}
|
||||
// @Router /env [post]
|
||||
func (ec *EnvController) CreateEnvVar(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Value string `json:"value" binding:"required"`
|
||||
Remark string `json:"remark"`
|
||||
Type string `json:"type"`
|
||||
Hidden *bool `json:"hidden"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
Tags string `json:"tags"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.Type == "" {
|
||||
req.Type = constant.EnvTypeNormal
|
||||
}
|
||||
|
||||
hidden := true
|
||||
if req.Hidden != nil {
|
||||
hidden = *req.Hidden
|
||||
}
|
||||
|
||||
enabled := true
|
||||
if req.Enabled != nil {
|
||||
enabled = *req.Enabled
|
||||
}
|
||||
|
||||
envVar := ec.envService.CreateEnvVar(req.Name, req.Value, req.Remark, req.Type, hidden, enabled, userID)
|
||||
if envVar != nil {
|
||||
relation.DataRelation.SaveTags(envVar.ID, constant.RelationTypeEnvTag, req.Tags)
|
||||
envVar.Tags = req.Tags
|
||||
}
|
||||
|
||||
// Broadcast tasks to all agents because global envs changed
|
||||
services.GetAgentWSManager().BroadcastTasksToAll()
|
||||
|
||||
utils.Success(c, vo.ToEnvVO(envVar))
|
||||
}
|
||||
|
||||
// GetEnvVars 获取环境变量列表
|
||||
// @Summary 获取环境变量列表
|
||||
// @Description 分页获取环境变量列表,支持按名称筛选
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param name query string false "按名称模糊查询"
|
||||
// @Param page query int false "页码"
|
||||
// @Param page_size query int false "每页数量"
|
||||
// @Param type query string false "按类型筛选"
|
||||
// @Param tags query string false "按标签筛选"
|
||||
// @Success 200 {object} utils.Response{data=utils.PaginationData{data=[]vo.EnvVO}}
|
||||
// @Router /env [get]
|
||||
func (ec *EnvController) GetEnvVars(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
p := utils.ParsePagination(c)
|
||||
name := c.DefaultQuery("name", "")
|
||||
envType := c.DefaultQuery("type", "")
|
||||
tags := c.DefaultQuery("tags", "")
|
||||
envVars, total := ec.envService.GetEnvVarsWithPagination(userID, name, envType, tags, p.Page, p.PageSize)
|
||||
utils.PaginatedResponse(c, vo.ToEnvVOListFromModels(envVars), total, p)
|
||||
}
|
||||
|
||||
// GetAllEnvVars 获取所有环境变量
|
||||
// @Summary 获取所有环境变量
|
||||
// @Description 获取当前用户的所有环境变量(不分页)
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Success 200 {object} utils.Response{data=[]vo.EnvVO}
|
||||
// @Router /env/all [get]
|
||||
func (ec *EnvController) GetAllEnvVars(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
envVars := ec.envService.GetEnvVarsByUserID(userID)
|
||||
utils.Success(c, vo.ToEnvVOListFromModels(envVars))
|
||||
}
|
||||
|
||||
// GetEnvVar 获取环境变量详情
|
||||
// @Summary 获取环境变量详情
|
||||
// @Description 根据 ID 获取环境变量详情
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "环境变量ID"
|
||||
// @Success 200 {object} utils.Response{data=vo.EnvVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /env/{id} [get]
|
||||
func (ec *EnvController) GetEnvVar(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的环境变量ID")
|
||||
return
|
||||
}
|
||||
|
||||
envVar := ec.envService.GetEnvVarByID(id)
|
||||
if envVar == nil {
|
||||
utils.NotFound(c, "环境变量不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToEnvVO(envVar))
|
||||
}
|
||||
|
||||
// UpdateEnvVar 更新环境变量
|
||||
// @Summary 更新环境变量
|
||||
// @Description 根据 ID 更新环境变量信息
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "环境变量ID"
|
||||
// @Param body body object true "环境变量更新信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.EnvVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /env/{id} [put]
|
||||
func (ec *EnvController) UpdateEnvVar(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的环境变量ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name"`
|
||||
Value string `json:"value"`
|
||||
Remark string `json:"remark"`
|
||||
Type string `json:"type"`
|
||||
Hidden *bool `json:"hidden"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
Tags string `json:"tags"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.Type == "" {
|
||||
req.Type = constant.EnvTypeNormal
|
||||
}
|
||||
|
||||
// 对于更新,获取现有数据
|
||||
existing := ec.envService.GetEnvVarByID(id)
|
||||
if existing == nil {
|
||||
utils.NotFound(c, "环境变量不存在")
|
||||
return
|
||||
}
|
||||
|
||||
hidden := existing.Hidden
|
||||
if req.Hidden != nil {
|
||||
hidden = req.Hidden
|
||||
}
|
||||
|
||||
enabled := existing.Enabled
|
||||
if req.Enabled != nil {
|
||||
enabled = req.Enabled
|
||||
}
|
||||
|
||||
envVar := ec.envService.UpdateEnvVar(id, req.Name, req.Value, req.Remark, req.Type, utils.DerefBool(hidden, true), utils.DerefBool(enabled, true))
|
||||
if envVar == nil {
|
||||
utils.NotFound(c, "环境变量不存在")
|
||||
return
|
||||
}
|
||||
|
||||
relation.DataRelation.SaveTags(envVar.ID, constant.RelationTypeEnvTag, req.Tags)
|
||||
envVar.Tags = req.Tags
|
||||
|
||||
// Broadcast tasks to all agents because global envs changed
|
||||
services.GetAgentWSManager().BroadcastTasksToAll()
|
||||
|
||||
utils.Success(c, vo.ToEnvVO(envVar))
|
||||
}
|
||||
|
||||
// DeleteEnvVar 删除环境变量
|
||||
// @Summary 删除环境变量
|
||||
// @Description 根据 ID 删除环境变量
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "环境变量ID"
|
||||
// @Param force query boolean false "强制删除(忽略任务关联)"
|
||||
// @Success 200 {object} utils.Response
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Failure 409 {object} utils.Response{data=[]vo.TaskVO}
|
||||
// @Router /env/{id} [delete]
|
||||
func (ec *EnvController) DeleteEnvVar(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的环境变量ID")
|
||||
return
|
||||
}
|
||||
|
||||
force := c.Query("force") == "true"
|
||||
success, associatedTasks := ec.envService.DeleteEnvVar(id, force)
|
||||
|
||||
if len(associatedTasks) > 0 {
|
||||
c.JSON(200, utils.Response{
|
||||
Code: 409,
|
||||
Msg: "该环境变量已被任务引用,请先在任务中删除引用或选择强制删除",
|
||||
Data: vo.ToTaskVOListFromModels(associatedTasks),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if !success {
|
||||
utils.NotFound(c, "环境变量不存在或删除失败")
|
||||
return
|
||||
}
|
||||
|
||||
// Broadcast tasks to all agents because global envs changed
|
||||
services.GetAgentWSManager().BroadcastTasksToAll()
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
// GetAssociatedTasks 获取关联任务
|
||||
// @Summary 获取关联任务
|
||||
// @Description 获取引用了该环境变量的任务列表
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "环境变量ID"
|
||||
// @Success 200 {object} utils.Response{data=[]vo.TaskVO}
|
||||
// @Router /env/{id}/tasks [get]
|
||||
func (ec *EnvController) GetAssociatedTasks(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的环境变量ID")
|
||||
return
|
||||
}
|
||||
tasks := ec.envService.GetAssociatedTasks(id)
|
||||
utils.Success(c, vo.ToTaskVOListFromModels(tasks))
|
||||
}
|
||||
|
||||
// GetTags 获取所有环境变量标签
|
||||
// @Summary 获取所有环境变量标签
|
||||
// @Description 获取所有环境变量中使用的标签列表
|
||||
// @Tags 环境变量
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Success 200 {object} utils.Response{data=[]string}
|
||||
// @Router /env/tags [get]
|
||||
func (ec *EnvController) GetTags(c *gin.Context) {
|
||||
tags, err := ec.envService.GetAllEnvTags()
|
||||
if err != nil {
|
||||
utils.ServerError(c, "获取标签失败")
|
||||
return
|
||||
}
|
||||
utils.Success(c, tags)
|
||||
}
|
||||
|
||||
// BulkSaveEnv 批量保存环境变量
|
||||
func (ec *EnvController) BulkSaveEnv(c *gin.Context) {
|
||||
var reqs []struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name" binding:"required"`
|
||||
Value string `json:"value" binding:"required"`
|
||||
Remark string `json:"remark"`
|
||||
Type string `json:"type"`
|
||||
Hidden *bool `json:"hidden"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&reqs); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
userID := c.GetString("userID")
|
||||
|
||||
for _, req := range reqs {
|
||||
if req.Type == constant.EnvTypeSecret {
|
||||
continue // 二次严苛拦截,机密变量不应下发/保存
|
||||
}
|
||||
|
||||
hidden := true
|
||||
if req.Hidden != nil {
|
||||
hidden = *req.Hidden
|
||||
}
|
||||
enabled := true
|
||||
if req.Enabled != nil {
|
||||
enabled = *req.Enabled
|
||||
}
|
||||
|
||||
var existingEnv *models.EnvironmentVariable
|
||||
// 优先按 ID 匹配
|
||||
if req.ID != "" {
|
||||
var e models.EnvironmentVariable
|
||||
if err := database.DB.Where("id = ?", req.ID).First(&e).Error; err == nil {
|
||||
existingEnv = &e
|
||||
}
|
||||
}
|
||||
// 如果 ID 没找到,按 Name 匹配
|
||||
if existingEnv == nil {
|
||||
var e models.EnvironmentVariable
|
||||
if err := database.DB.Where("name = ?", req.Name).First(&e).Error; err == nil {
|
||||
existingEnv = &e
|
||||
}
|
||||
}
|
||||
|
||||
if existingEnv != nil {
|
||||
existingEnv.Name = req.Name
|
||||
existingEnv.Value = models.BigText(req.Value)
|
||||
existingEnv.Remark = req.Remark
|
||||
existingEnv.Type = req.Type
|
||||
existingEnv.Hidden = &hidden
|
||||
existingEnv.Enabled = &enabled
|
||||
database.DB.Save(existingEnv)
|
||||
|
||||
if req.ID != "" && existingEnv.ID != req.ID {
|
||||
database.DB.Model(existingEnv).Update("id", req.ID)
|
||||
}
|
||||
} else {
|
||||
envVar := &models.EnvironmentVariable{
|
||||
ID: req.ID,
|
||||
Name: req.Name,
|
||||
Value: models.BigText(req.Value),
|
||||
Remark: req.Remark,
|
||||
Type: req.Type,
|
||||
Hidden: &hidden,
|
||||
Enabled: &enabled,
|
||||
UserID: userID,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
if envVar.ID == "" {
|
||||
envVar.ID = utils.GenerateID()
|
||||
}
|
||||
database.DB.Create(envVar)
|
||||
}
|
||||
}
|
||||
|
||||
services.GetAgentWSManager().BroadcastTasksToAll()
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type ExecutorController struct {
|
||||
executorService *tasks.ExecutorService
|
||||
}
|
||||
|
||||
func NewExecutorController(executorService *tasks.ExecutorService) *ExecutorController {
|
||||
return &ExecutorController{executorService: executorService}
|
||||
}
|
||||
|
||||
// ExecuteTask 运行任务
|
||||
// @Summary 运行任务
|
||||
// @Description 立即执行指定的任务
|
||||
// @Tags 任务执行
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "任务ID"
|
||||
// @Param body body object false "执行参数 (envs: 环境变量字典)"
|
||||
// @Success 200 {object} utils.Response{data=vo.ExecutionResultVO}
|
||||
// @Failure 400 {object} utils.Response
|
||||
// @Router /execute/task/{id} [post]
|
||||
func (ec *ExecutorController) ExecuteTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Envs map[string]string `json:"envs"`
|
||||
}
|
||||
// 尝试绑定 JSON 体,但不强制要求
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
|
||||
var extraEnvs []string
|
||||
if req.Envs != nil {
|
||||
for k, v := range req.Envs {
|
||||
extraEnvs = append(extraEnvs, k+"="+v)
|
||||
}
|
||||
}
|
||||
|
||||
result := ec.executorService.ExecuteTask(id, extraEnvs)
|
||||
utils.Success(c, vo.ToExecutionResultVO(result))
|
||||
}
|
||||
|
||||
// ExecuteCommand 执行命令
|
||||
func (ec *ExecutorController) ExecuteCommand(c *gin.Context) {
|
||||
var req struct {
|
||||
Command string `json:"command" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
result := ec.executorService.ExecuteCommand(req.Command)
|
||||
utils.Success(c, vo.ToExecutionResultVO(result))
|
||||
}
|
||||
|
||||
// GetLastResults 获取最新执行结果
|
||||
// @Summary 获取最新执行结果
|
||||
// @Description 获取最新任务或命令执行的结果列表
|
||||
// @Tags 任务执行
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param count query int false "数量 (默认 10)"
|
||||
// @Success 200 {object} utils.Response{data=[]vo.ExecutionResultVO}
|
||||
// @Router /execute/results [get]
|
||||
func (ec *ExecutorController) GetLastResults(c *gin.Context) {
|
||||
count := 10
|
||||
if c.Query("count") != "" {
|
||||
if parsedCount, err := strconv.Atoi(c.Query("count")); err == nil && parsedCount > 0 {
|
||||
count = parsedCount
|
||||
}
|
||||
}
|
||||
|
||||
results := ec.executorService.GetLastResults(count)
|
||||
utils.Success(c, vo.ToExecutionResultVOList(results))
|
||||
}
|
||||
@@ -0,0 +1,550 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
var (
|
||||
extractZip = utils.ExtractZip
|
||||
extractTar = utils.ExtractTar
|
||||
extractTarGz = utils.ExtractTarGz
|
||||
)
|
||||
|
||||
type FileController struct {
|
||||
workDir string
|
||||
}
|
||||
|
||||
func NewFileController(workDir string) *FileController {
|
||||
os.MkdirAll(workDir, 0755)
|
||||
absPath, err := filepath.Abs(workDir)
|
||||
if err != nil {
|
||||
absPath = workDir
|
||||
}
|
||||
return &FileController{workDir: absPath}
|
||||
}
|
||||
|
||||
type FileNode struct {
|
||||
Name string `json:"name"`
|
||||
Path string `json:"path"`
|
||||
IsDir bool `json:"isDir"`
|
||||
ModTime int64 `json:"modTime"`
|
||||
Children []*FileNode `json:"children,omitempty"`
|
||||
}
|
||||
|
||||
// checkPath 校验路径是否在工作目录内且安全。
|
||||
// 它返回完整的绝对路径以及一个表示路径是否安全的布尔值。
|
||||
func (fc *FileController) checkPath(path string, allowRoot bool) (string, bool) {
|
||||
fullPath := filepath.Join(fc.workDir, filepath.Clean(path))
|
||||
rel, err := filepath.Rel(fc.workDir, fullPath)
|
||||
if err != nil {
|
||||
return "", false
|
||||
}
|
||||
|
||||
// 基础的目录穿越检查
|
||||
if rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
|
||||
return "", false
|
||||
}
|
||||
|
||||
// 根目录检查
|
||||
if !allowRoot && rel == "." {
|
||||
return "", false
|
||||
}
|
||||
|
||||
return fullPath, true
|
||||
}
|
||||
|
||||
func (fc *FileController) GetFileTree(c *gin.Context) {
|
||||
root := &FileNode{
|
||||
Name: filepath.Base(fc.workDir),
|
||||
Path: "",
|
||||
IsDir: true,
|
||||
Children: []*FileNode{},
|
||||
}
|
||||
|
||||
err := filepath.WalkDir(fc.workDir, func(path string, d fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
if path == fc.workDir {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 过滤 __pycache__ 文件夹
|
||||
if d.IsDir() && d.Name() == "__pycache__" {
|
||||
return filepath.SkipDir
|
||||
}
|
||||
|
||||
relPath, _ := filepath.Rel(fc.workDir, path)
|
||||
parts := strings.Split(relPath, string(filepath.Separator))
|
||||
|
||||
info, err := d.Info()
|
||||
var modTime int64
|
||||
if err == nil {
|
||||
modTime = info.ModTime().UnixMilli()
|
||||
}
|
||||
|
||||
current := root
|
||||
for i, part := range parts {
|
||||
found := false
|
||||
for _, child := range current.Children {
|
||||
if child.Name == part {
|
||||
current = child
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
isLast := i == len(parts)-1
|
||||
isDir := !isLast || d.IsDir()
|
||||
node := &FileNode{
|
||||
Name: part,
|
||||
Path: strings.Join(parts[:i+1], "/"),
|
||||
IsDir: isDir,
|
||||
ModTime: modTime,
|
||||
}
|
||||
if isDir {
|
||||
node.Children = []*FileNode{}
|
||||
}
|
||||
current.Children = append(current.Children, node)
|
||||
current = node
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, root.Children)
|
||||
}
|
||||
|
||||
func (fc *FileController) GetFileContent(c *gin.Context) {
|
||||
filePath := c.Query("path")
|
||||
if filePath == "" {
|
||||
utils.BadRequest(c, "path参数必填")
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(filePath, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
content, err := os.ReadFile(fullPath)
|
||||
if err != nil {
|
||||
utils.NotFound(c, "文件不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"path": filePath,
|
||||
"content": string(content),
|
||||
})
|
||||
}
|
||||
|
||||
func (fc *FileController) SaveFileContent(c *gin.Context) {
|
||||
var req struct {
|
||||
Path string `json:"path" binding:"required"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(req.Path, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
os.MkdirAll(filepath.Dir(fullPath), 0755)
|
||||
|
||||
if err := os.WriteFile(fullPath, []byte(req.Content), 0644); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
func (fc *FileController) CreateFile(c *gin.Context) {
|
||||
var req struct {
|
||||
Path string `json:"path" binding:"required"`
|
||||
IsDir bool `json:"isDir"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(req.Path, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if req.IsDir {
|
||||
if err := os.MkdirAll(fullPath, 0755); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
} else {
|
||||
os.MkdirAll(filepath.Dir(fullPath), 0755)
|
||||
if err := os.WriteFile(fullPath, []byte(""), 0644); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "创建成功")
|
||||
}
|
||||
|
||||
func (fc *FileController) DeleteFile(c *gin.Context) {
|
||||
var req struct {
|
||||
Path string `json:"path" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(req.Path, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if err := os.RemoveAll(fullPath); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
func (fc *FileController) MoveFile(c *gin.Context) {
|
||||
var req struct {
|
||||
OldPath string `json:"oldPath" binding:"required"`
|
||||
NewPath string `json:"newPath" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
oldFull, oldSafe := fc.checkPath(req.OldPath, false)
|
||||
newFull, newSafe := fc.checkPath(req.NewPath, false)
|
||||
|
||||
if !oldSafe || !newSafe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if oldFull == newFull {
|
||||
utils.Success(c, nil)
|
||||
return
|
||||
}
|
||||
|
||||
// 检查目标是否存在
|
||||
if _, err := os.Stat(newFull); err == nil {
|
||||
utils.BadRequest(c, "目标已存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 确保目标目录存在
|
||||
os.MkdirAll(filepath.Dir(newFull), 0755)
|
||||
|
||||
if err := os.Rename(oldFull, newFull); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
func (fc *FileController) CopyFile(c *gin.Context) {
|
||||
var req struct {
|
||||
SourcePath string `json:"sourcePath" binding:"required"`
|
||||
TargetPath string `json:"targetPath" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
sourceFull, sourceSafe := fc.checkPath(req.SourcePath, false)
|
||||
targetFull, targetSafe := fc.checkPath(req.TargetPath, false)
|
||||
|
||||
if !sourceSafe || !targetSafe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if sourceFull == targetFull {
|
||||
utils.Success(c, nil)
|
||||
return
|
||||
}
|
||||
|
||||
// Read content
|
||||
content, err := os.ReadFile(sourceFull)
|
||||
if err != nil {
|
||||
utils.NotFound(c, "源文件不存在或无法读取")
|
||||
return
|
||||
}
|
||||
|
||||
// 确保目标目录存在
|
||||
os.MkdirAll(filepath.Dir(targetFull), 0755)
|
||||
|
||||
// 检查目标是否存在
|
||||
if _, err := os.Stat(targetFull); err == nil {
|
||||
utils.BadRequest(c, "目标已存在")
|
||||
return
|
||||
}
|
||||
|
||||
if err := os.WriteFile(targetFull, content, 0644); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
func (fc *FileController) RenameFile(c *gin.Context) {
|
||||
var req struct {
|
||||
OldPath string `json:"oldPath" binding:"required"`
|
||||
NewPath string `json:"newPath" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 校验:重命名禁止跨目录
|
||||
if filepath.Dir(filepath.Clean(req.OldPath)) != filepath.Dir(filepath.Clean(req.NewPath)) {
|
||||
utils.BadRequest(c, "禁止跨目录重命名")
|
||||
return
|
||||
}
|
||||
|
||||
oldFull, oldSafe := fc.checkPath(req.OldPath, false)
|
||||
newFull, newSafe := fc.checkPath(req.NewPath, false)
|
||||
|
||||
if !oldSafe || !newSafe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if oldFull == newFull {
|
||||
utils.Success(c, nil)
|
||||
return
|
||||
}
|
||||
|
||||
// 检查目标是否存在
|
||||
if _, err := os.Stat(newFull); err == nil {
|
||||
utils.BadRequest(c, "文件已存在")
|
||||
return
|
||||
}
|
||||
|
||||
if err := os.Rename(oldFull, newFull); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
// UploadArchive 处理归档文件的上传和解压
|
||||
func (fc *FileController) UploadArchive(c *gin.Context) {
|
||||
targetDir := c.PostForm("path")
|
||||
|
||||
file, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
utils.BadRequest(c, "请选择文件")
|
||||
return
|
||||
}
|
||||
|
||||
// 检查文件类型
|
||||
ext := strings.ToLower(filepath.Ext(file.Filename))
|
||||
if ext != ".zip" && ext != ".tar" && ext != ".gz" && ext != ".tgz" {
|
||||
utils.BadRequest(c, "仅支持 zip、tar、gz、tgz 格式")
|
||||
return
|
||||
}
|
||||
|
||||
// 确定解压目标目录
|
||||
extractDir, safe := fc.checkPath(targetDir, true)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
os.MkdirAll(extractDir, 0755)
|
||||
|
||||
// 保存临时文件
|
||||
// 安全修复:使用 filepath.Base 提取纯文件名,防止路径穿越攻击
|
||||
tempFile := filepath.Join(os.TempDir(), filepath.Base(file.Filename))
|
||||
if err := c.SaveUploadedFile(file, tempFile); err != nil {
|
||||
utils.ServerError(c, "保存文件失败")
|
||||
return
|
||||
}
|
||||
defer os.Remove(tempFile)
|
||||
|
||||
// 解压文件
|
||||
var extractErr error
|
||||
switch {
|
||||
case ext == ".zip":
|
||||
extractErr = extractZip(tempFile, extractDir)
|
||||
case ext == ".tar":
|
||||
extractErr = extractTar(tempFile, extractDir)
|
||||
case ext == ".gz" || ext == ".tgz":
|
||||
extractErr = extractTarGz(tempFile, extractDir)
|
||||
}
|
||||
|
||||
if extractErr != nil {
|
||||
utils.ServerError(c, "解压失败: "+extractErr.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "导入成功")
|
||||
}
|
||||
|
||||
// UploadFiles 处理多个文件的上传
|
||||
func (fc *FileController) UploadFiles(c *gin.Context) {
|
||||
targetDir := c.PostForm("path")
|
||||
|
||||
// 确定目标目录
|
||||
destDir, safe := fc.checkPath(targetDir, true)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
os.MkdirAll(destDir, 0755)
|
||||
|
||||
form, err := c.MultipartForm()
|
||||
if err != nil {
|
||||
utils.BadRequest(c, "请选择文件")
|
||||
return
|
||||
}
|
||||
|
||||
files := form.File["files"]
|
||||
paths := form.Value["paths"] // 相对路径数组,用于保持文件夹结构
|
||||
|
||||
if len(files) == 0 {
|
||||
utils.BadRequest(c, "请选择文件")
|
||||
return
|
||||
}
|
||||
|
||||
for i, file := range files {
|
||||
// 获取相对路径(如果有)
|
||||
// 安全修复:清理文件名
|
||||
relPath := filepath.Base(file.Filename)
|
||||
if i < len(paths) && paths[i] != "" {
|
||||
relPath = paths[i]
|
||||
}
|
||||
|
||||
// 构建完整路径
|
||||
fullPath, safe := fc.checkPath(filepath.Join(targetDir, relPath), false)
|
||||
if !safe {
|
||||
continue
|
||||
}
|
||||
|
||||
// 确保父目录存在
|
||||
os.MkdirAll(filepath.Dir(fullPath), 0755)
|
||||
|
||||
// 保存文件
|
||||
if err := c.SaveUploadedFile(file, fullPath); err != nil {
|
||||
utils.ServerError(c, "保存文件失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "上传成功")
|
||||
}
|
||||
|
||||
func (fc *FileController) DownloadFile(c *gin.Context) {
|
||||
filePath := c.Query("path")
|
||||
if filePath == "" {
|
||||
utils.BadRequest(c, "path参数必填")
|
||||
return
|
||||
}
|
||||
|
||||
fullPath, safe := fc.checkPath(filePath, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
info, err := os.Stat(fullPath)
|
||||
if err != nil || info.IsDir() {
|
||||
utils.NotFound(c, "文件不存在")
|
||||
return
|
||||
}
|
||||
|
||||
c.Header("Content-Description", "File Transfer")
|
||||
c.Header("Content-Transfer-Encoding", "binary")
|
||||
c.Header("Content-Disposition", "attachment; filename="+filepath.Base(fullPath))
|
||||
c.Header("Content-Type", "application/octet-stream")
|
||||
c.File(fullPath)
|
||||
}
|
||||
|
||||
func (fc *FileController) DownloadZip(c *gin.Context) {
|
||||
paths := c.QueryArray("path")
|
||||
if len(paths) == 0 || c.ContentType() == "application/json" {
|
||||
var req struct {
|
||||
Paths []string `json:"paths"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err == nil && len(paths) == 0 {
|
||||
paths = req.Paths
|
||||
}
|
||||
}
|
||||
|
||||
if len(paths) == 0 {
|
||||
utils.BadRequest(c, "path参数必填")
|
||||
return
|
||||
}
|
||||
|
||||
validatedAbsPaths := make([]string, 0, len(paths))
|
||||
for _, path := range paths {
|
||||
fullPath, safe := fc.checkPath(path, false)
|
||||
if !safe {
|
||||
utils.Forbidden(c, "访问被拒绝")
|
||||
return
|
||||
}
|
||||
|
||||
if _, err := os.Stat(fullPath); err != nil {
|
||||
utils.NotFound(c, "文件不存在")
|
||||
return
|
||||
}
|
||||
|
||||
validatedAbsPaths = append(validatedAbsPaths, fullPath)
|
||||
}
|
||||
|
||||
fileName := "taskpool-export-" + time.Now().Format("20060102-150405") + ".zip"
|
||||
if len(validatedAbsPaths) == 1 {
|
||||
fileName = filepath.Base(validatedAbsPaths[0]) + ".zip"
|
||||
}
|
||||
|
||||
c.Header("Content-Description", "File Transfer")
|
||||
c.Header("Content-Transfer-Encoding", "binary")
|
||||
c.Header("Content-Disposition", "attachment; filename="+fileName)
|
||||
c.Header("Content-Type", "application/zip")
|
||||
|
||||
if err := utils.CreateZip(c.Writer, validatedAbsPaths); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type InstallController struct {
|
||||
installService *services.InstallService
|
||||
}
|
||||
|
||||
func NewInstallController() *InstallController {
|
||||
return &InstallController{
|
||||
installService: services.NewInstallService(),
|
||||
}
|
||||
}
|
||||
|
||||
// GetInstallStatus 获取安装状态
|
||||
// @Summary 获取安装状态
|
||||
// @Description 检查系统是否已完成安装
|
||||
// @Tags 安装
|
||||
// @Produce json
|
||||
// @Success 200 {object} services.InstallStatus
|
||||
// @Router /api/v1/install/status [get]
|
||||
func (c *InstallController) GetInstallStatus(ctx *gin.Context) {
|
||||
status, err := c.installService.CheckInstallStatus()
|
||||
if err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
ctx.JSON(http.StatusOK, status)
|
||||
}
|
||||
|
||||
// Install 执行安装
|
||||
// @Summary 执行安装
|
||||
// @Description 初始化系统配置和管理员账号
|
||||
// @Tags 安装
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param request body services.InstallRequest true "安装请求"
|
||||
// @Success 200 {object} map[string]interface{}
|
||||
// @Failure 400 {object} map[string]string
|
||||
// @Failure 500 {object} map[string]string
|
||||
// @Router /api/v1/install [post]
|
||||
func (c *InstallController) Install(ctx *gin.Context) {
|
||||
// 先检查是否已安装
|
||||
status, err := c.installService.CheckInstallStatus()
|
||||
if err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, gin.H{"error": "检查安装状态失败"})
|
||||
return
|
||||
}
|
||||
|
||||
if status.Installed {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "系统已安装,无法重复安装"})
|
||||
return
|
||||
}
|
||||
|
||||
var req services.InstallRequest
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "参数错误: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
// 验证必填字段
|
||||
if req.AdminUsername == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "管理员用户名不能为空"})
|
||||
return
|
||||
}
|
||||
if req.AdminPassword == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "管理员密码不能为空"})
|
||||
return
|
||||
}
|
||||
if len(req.AdminPassword) < 6 {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "管理员密码至少6位"})
|
||||
return
|
||||
}
|
||||
|
||||
// MySQL 必填验证
|
||||
if req.DBType == "mysql" {
|
||||
if req.DBHost == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "MySQL 主机不能为空"})
|
||||
return
|
||||
}
|
||||
if req.DBName == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "MySQL 数据库名不能为空"})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Redis 启用时的验证
|
||||
if req.RedisEnabled {
|
||||
if req.RedisHost == "" {
|
||||
ctx.JSON(http.StatusBadRequest, gin.H{"error": "Redis 主机不能为空"})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 执行安装
|
||||
if err := c.installService.Install(&req); err != nil {
|
||||
ctx.JSON(http.StatusInternalServerError, gin.H{"error": "安装失败: " + err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
ctx.JSON(http.StatusOK, gin.H{
|
||||
"message": "安装成功",
|
||||
"admin_username": req.AdminUsername,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,532 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/tunnel"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type InterconnectController struct {
|
||||
interconnectService *services.InterconnectService
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewInterconnectController(interconnectService *services.InterconnectService) *InterconnectController {
|
||||
return &InterconnectController{
|
||||
interconnectService: interconnectService,
|
||||
httpClient: &http.Client{
|
||||
Timeout: 10 * time.Second,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// GetNodes 获取互联节点列表
|
||||
func (ic *InterconnectController) GetNodes(c *gin.Context) {
|
||||
nodes, err := ic.interconnectService.GetNodes()
|
||||
if err != nil {
|
||||
utils.ServerError(c, "获取互联节点失败")
|
||||
return
|
||||
}
|
||||
utils.Success(c, nodes)
|
||||
}
|
||||
|
||||
// CreateNode 创建互联节点
|
||||
func (ic *InterconnectController) CreateNode(c *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
URL string `json:"url"`
|
||||
Token string `json:"token" binding:"required"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
node, err := ic.interconnectService.CreateNode(req.Name, req.URL, req.Token, req.Remark)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "创建互联节点失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, node)
|
||||
}
|
||||
|
||||
// UpdateNode 更新互联节点
|
||||
func (ic *InterconnectController) UpdateNode(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的节点ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
URL string `json:"url"`
|
||||
Token string `json:"token" binding:"required"`
|
||||
Remark string `json:"remark"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
node, err := ic.interconnectService.UpdateNode(id, req.Name, req.URL, req.Token, req.Remark)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "更新互联节点失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, node)
|
||||
}
|
||||
|
||||
// DeleteNode 删除互联节点
|
||||
func (ic *InterconnectController) DeleteNode(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的节点ID")
|
||||
return
|
||||
}
|
||||
|
||||
err := ic.interconnectService.DeleteNode(id)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "删除互联节点失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
// GetNodeStatus 获取单个子节点的状态
|
||||
func (ic *InterconnectController) GetNodeStatus(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
node, err := ic.interconnectService.GetNodeByID(id)
|
||||
if err != nil {
|
||||
utils.NotFound(c, "节点不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 针对反向隧道节点状态检测的特判
|
||||
if strings.HasPrefix(node.URL, "tunnel://") {
|
||||
sess := tunnel.GetSession(node.ID)
|
||||
if sess == nil {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "节点离线或反向隧道未建立", "data": nil})
|
||||
return
|
||||
}
|
||||
|
||||
// 使用当前 Yamux Session 的虚拟底层连接进行拨号
|
||||
transport := &http.Transport{
|
||||
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
return sess.Session.Open()
|
||||
},
|
||||
}
|
||||
client := &http.Client{
|
||||
Transport: transport,
|
||||
Timeout: 5 * time.Second,
|
||||
}
|
||||
|
||||
req, err := http.NewRequest("GET", "http://tunnel.local/api/v1/monitor", nil)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "构建检测请求失败")
|
||||
return
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "与子节点逆向连接通讯失败", "data": nil})
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode != 200 {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "子节点检测异常", "data": string(body)})
|
||||
return
|
||||
}
|
||||
|
||||
var jsonResp map[string]interface{}
|
||||
if err := json.Unmarshal(body, &jsonResp); err != nil {
|
||||
utils.ServerError(c, "解析节点检测数据失败")
|
||||
return
|
||||
}
|
||||
|
||||
if dataMap, ok := jsonResp["data"].(map[string]interface{}); ok {
|
||||
dataMap["tunnel_connected"] = true
|
||||
dataMap["tunnel_url"] = node.URL
|
||||
if hostMap, ok := dataMap["host"].(map[string]interface{}); ok {
|
||||
hostMap["tx_bytes"] = node.Metrics.TxBytes
|
||||
hostMap["rx_bytes"] = node.Metrics.RxBytes
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, jsonResp["data"])
|
||||
return
|
||||
}
|
||||
|
||||
apiURL := strings.TrimRight(node.URL, "/") + "/api/v1/monitor"
|
||||
req, err := http.NewRequest("GET", apiURL, nil)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "构建请求失败")
|
||||
return
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
|
||||
resp, err := ic.httpClient.Do(req)
|
||||
if err != nil {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "节点离线或网络不可达", "data": nil})
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
if resp.StatusCode != 200 {
|
||||
c.JSON(200, gin.H{"code": 500, "msg": "节点返回异常", "data": string(body)})
|
||||
return
|
||||
}
|
||||
|
||||
var jsonResp map[string]interface{}
|
||||
if err := json.Unmarshal(body, &jsonResp); err != nil {
|
||||
utils.ServerError(c, "解析节点响应失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, jsonResp["data"])
|
||||
}
|
||||
|
||||
// SyncScript 将脚本同步到指定的节点列表
|
||||
func (ic *InterconnectController) SyncScript(c *gin.Context) {
|
||||
var req struct {
|
||||
NodeIDs []string `json:"node_ids" binding:"required"`
|
||||
Filename string `json:"filename" binding:"required"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
results := make([]map[string]interface{}, 0)
|
||||
|
||||
for _, nodeID := range req.NodeIDs {
|
||||
node, err := ic.interconnectService.GetNodeByID(nodeID)
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "节点不存在"})
|
||||
continue
|
||||
}
|
||||
|
||||
client, apiURL, err := ic.getClientAndURL(node, "/api/v1/scripts/save")
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "反向隧道未连接"})
|
||||
continue
|
||||
}
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"filename": req.Filename,
|
||||
"content": req.Content,
|
||||
}
|
||||
payloadBytes, _ := json.Marshal(payload)
|
||||
|
||||
httpReq, err := http.NewRequest("POST", apiURL, bytes.NewBuffer(payloadBytes))
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "构建请求失败"})
|
||||
continue
|
||||
}
|
||||
httpReq.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil || resp.StatusCode != 200 {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "同步请求失败或超时"})
|
||||
} else {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": true, "msg": "同步成功"})
|
||||
}
|
||||
if resp != nil {
|
||||
resp.Body.Close()
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, results)
|
||||
}
|
||||
|
||||
// SyncEnv 将环境变量同步到指定的节点列表
|
||||
func (ic *InterconnectController) SyncEnv(c *gin.Context) {
|
||||
var req struct {
|
||||
NodeIDs []string `json:"node_ids" binding:"required"`
|
||||
Envs []struct{ ID string `json:"id"` } `json:"envs" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var envIDs []string
|
||||
for _, e := range req.Envs {
|
||||
envIDs = append(envIDs, e.ID)
|
||||
}
|
||||
|
||||
dataService := services.NewDataService()
|
||||
exportData := dataService.ExportBusinessData(nil, envIDs)
|
||||
|
||||
results := make([]map[string]interface{}, 0)
|
||||
|
||||
for _, nodeID := range req.NodeIDs {
|
||||
node, err := ic.interconnectService.GetNodeByID(nodeID)
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "节点不存在"})
|
||||
continue
|
||||
}
|
||||
|
||||
client, apiURL, err := ic.getClientAndURL(node, "/api/v1/system/import")
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "反向隧道未连接"})
|
||||
continue
|
||||
}
|
||||
|
||||
payloadBytes, _ := json.Marshal(exportData)
|
||||
|
||||
httpReq, err := http.NewRequest("POST", apiURL, bytes.NewBuffer(payloadBytes))
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "构建请求失败"})
|
||||
continue
|
||||
}
|
||||
httpReq.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil || resp.StatusCode != 200 {
|
||||
msg := "同步失败"
|
||||
if err != nil {
|
||||
msg = err.Error()
|
||||
}
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": msg})
|
||||
} else {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": true, "msg": "同步成功"})
|
||||
}
|
||||
if resp != nil {
|
||||
resp.Body.Close()
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, results)
|
||||
}
|
||||
|
||||
// SyncTask 将任务同步到指定的节点列表
|
||||
func (ic *InterconnectController) SyncTask(c *gin.Context) {
|
||||
var req struct {
|
||||
NodeIDs []string `json:"node_ids" binding:"required"`
|
||||
Tasks []struct{ ID string `json:"id"` } `json:"tasks" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
var taskIDs []string
|
||||
for _, t := range req.Tasks {
|
||||
taskIDs = append(taskIDs, t.ID)
|
||||
}
|
||||
|
||||
dataService := services.NewDataService()
|
||||
exportData := dataService.ExportBusinessData(taskIDs, nil)
|
||||
|
||||
results := make([]map[string]interface{}, 0)
|
||||
|
||||
for _, nodeID := range req.NodeIDs {
|
||||
node, err := ic.interconnectService.GetNodeByID(nodeID)
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "节点不存在"})
|
||||
continue
|
||||
}
|
||||
|
||||
client, apiURL, err := ic.getClientAndURL(node, "/api/v1/system/import")
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "反向隧道未连接"})
|
||||
continue
|
||||
}
|
||||
|
||||
payloadBytes, _ := json.Marshal(exportData)
|
||||
|
||||
httpReq, err := http.NewRequest("POST", apiURL, bytes.NewBuffer(payloadBytes))
|
||||
if err != nil {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": "构建请求失败"})
|
||||
continue
|
||||
}
|
||||
httpReq.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
|
||||
resp, err := client.Do(httpReq)
|
||||
if err != nil || resp.StatusCode != 200 {
|
||||
msg := "同步失败"
|
||||
if err != nil {
|
||||
msg = err.Error()
|
||||
}
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": false, "msg": msg})
|
||||
} else {
|
||||
results = append(results, map[string]interface{}{"node_id": nodeID, "success": true, "msg": "同步成功"})
|
||||
}
|
||||
if resp != nil {
|
||||
resp.Body.Close()
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, results)
|
||||
}
|
||||
|
||||
// HandleTunnel 接受子节点 WebSocket 连接请求
|
||||
func (ic *InterconnectController) HandleTunnel(c *gin.Context) {
|
||||
tunnel.HandleTunnel(c)
|
||||
}
|
||||
|
||||
// ProxyRequest 代理转发请求至目标节点
|
||||
func (ic *InterconnectController) ProxyRequest(c *gin.Context) {
|
||||
nodeID := c.Param("node_id")
|
||||
path := c.Param("path")
|
||||
if nodeID == "" {
|
||||
utils.BadRequest(c, "Node ID required")
|
||||
return
|
||||
}
|
||||
|
||||
node, err := ic.interconnectService.GetNodeByID(nodeID)
|
||||
if err != nil {
|
||||
utils.NotFound(c, "Node not found")
|
||||
return
|
||||
}
|
||||
|
||||
if strings.HasPrefix(node.URL, "tunnel://") {
|
||||
// 走 WebSocket 逆向隧道 (基于 Yamux 流式多路复用)
|
||||
err := tunnel.ProxyHTTP(nodeID, c, path)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "Tunnel request failed: "+err.Error())
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// 走普通 HTTP 直连
|
||||
// Construct the target URL
|
||||
targetURL := strings.TrimRight(node.URL, "/") + path
|
||||
if c.Request.URL.RawQuery != "" {
|
||||
targetURL += "?" + c.Request.URL.RawQuery
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(c.Request.Method, targetURL, c.Request.Body)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "Failed to create proxy request")
|
||||
return
|
||||
}
|
||||
|
||||
// Copy headers
|
||||
req.Header = c.Request.Header.Clone()
|
||||
|
||||
// If the node token exists, append it as Bearer Auth
|
||||
if node.Token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
}
|
||||
|
||||
resp, err := ic.httpClient.Do(req)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "Failed to connect to target node: "+err.Error())
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
for k, v := range resp.Header {
|
||||
for _, vv := range v {
|
||||
c.Writer.Header().Add(k, vv)
|
||||
}
|
||||
}
|
||||
c.Status(resp.StatusCode)
|
||||
io.Copy(c.Writer, resp.Body)
|
||||
}
|
||||
|
||||
// ReportMonitorData 接收子节点上报的监控数据
|
||||
func (ic *InterconnectController) ReportMonitorData(c *gin.Context) {
|
||||
var req models.NodeMetrics
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if authHeader == "" {
|
||||
c.JSON(401, gin.H{"error": "missing authorization"})
|
||||
return
|
||||
}
|
||||
tokenStr := strings.TrimSpace(strings.TrimPrefix(authHeader, "Bearer "))
|
||||
|
||||
node, err := ic.interconnectService.GetNodeByToken(tokenStr)
|
||||
if err != nil {
|
||||
c.JSON(401, gin.H{"error": "invalid token"})
|
||||
return
|
||||
}
|
||||
|
||||
err = ic.interconnectService.UpdateNodeMonitorData(node.ID, req)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "更新节点数据失败")
|
||||
return
|
||||
}
|
||||
utils.Success(c, gin.H{
|
||||
"tunnel_url": node.URL,
|
||||
})
|
||||
}
|
||||
|
||||
// GetChildStatus 获取本机作为子节点的连接状态
|
||||
func (ic *InterconnectController) GetChildStatus(c *gin.Context) {
|
||||
settingsSvc := services.NewSettingsService()
|
||||
parentURL := settingsSvc.Get(constant.SectionInterconnect, constant.KeyInterconnectParentURL)
|
||||
parentToken := settingsSvc.Get(constant.SectionInterconnect, constant.KeyInterconnectParentToken)
|
||||
|
||||
connected := tunnel.IsTunnelConnected()
|
||||
tunnelURL := tunnel.GetLocalTunnelURL()
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"parent_url": parentURL,
|
||||
"parent_token": parentToken,
|
||||
"connected": connected,
|
||||
"tunnel_url": tunnelURL,
|
||||
"tx_bytes": tunnel.GetTxBytes(),
|
||||
"rx_bytes": tunnel.GetRxBytes(),
|
||||
})
|
||||
}
|
||||
|
||||
// getClientAndURL 辅助方法:根据节点类型决定走直连还是隧道,并返回对应的 Client 和完整 URL
|
||||
func (ic *InterconnectController) getClientAndURL(node *models.InterconnectNode, path string) (*http.Client, string, error) {
|
||||
if strings.HasPrefix(node.URL, "tunnel://") {
|
||||
sess := tunnel.GetSession(node.ID)
|
||||
if sess == nil {
|
||||
return nil, "", net.ErrClosed
|
||||
}
|
||||
transport := &http.Transport{
|
||||
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
return sess.Session.Open()
|
||||
},
|
||||
}
|
||||
client := &http.Client{
|
||||
Transport: transport,
|
||||
Timeout: 10 * time.Second,
|
||||
}
|
||||
return client, "http://tunnel.local" + path, nil
|
||||
}
|
||||
|
||||
targetURL := strings.TrimRight(node.URL, "/") + path
|
||||
return ic.httpClient, targetURL, nil
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type LogController struct{}
|
||||
|
||||
func NewLogController() *LogController {
|
||||
return &LogController{}
|
||||
}
|
||||
|
||||
// GetLogs 获取任务日志列表
|
||||
// @Summary 获取任务日志列表
|
||||
// @Description 分页获取任务日志列表,支持按任务 ID、任务名称、状态筛选
|
||||
// @Tags 日志管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param task_id query string false "任务 ID"
|
||||
// @Param task_name query string false "任务名称"
|
||||
// @Param status query string false "状态"
|
||||
// @Param page query int false "页码"
|
||||
// @Param page_size query int false "每页数量"
|
||||
// @Success 200 {object} utils.Response{data=utils.PaginationData{data=[]vo.TaskLogVO}}
|
||||
// @Router /logs [get]
|
||||
func (lc *LogController) GetLogs(c *gin.Context) {
|
||||
p := utils.ParsePagination(c)
|
||||
taskID := c.DefaultQuery("task_id", "")
|
||||
taskName := c.DefaultQuery("task_name", "")
|
||||
status := c.DefaultQuery("status", "")
|
||||
|
||||
var logs []models.TaskLog
|
||||
var total int64
|
||||
|
||||
query := database.DB.Model(&models.TaskLog{})
|
||||
if taskID != "" {
|
||||
query = query.Where("task_id = ?", taskID)
|
||||
}
|
||||
if status != "" {
|
||||
query = query.Where("status = ?", status)
|
||||
}
|
||||
|
||||
// 按任务名称过滤
|
||||
if taskName != "" {
|
||||
var taskIDs []string
|
||||
database.DB.Model(&models.Task{}).Where("name LIKE ?", "%"+taskName+"%").Pluck("id", &taskIDs)
|
||||
if len(taskIDs) > 0 {
|
||||
query = query.Where("task_id IN ?", taskIDs)
|
||||
} else {
|
||||
utils.PaginatedResponse(c, []vo.TaskLogVO{}, 0, p)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
query.Count(&total)
|
||||
query.Order("id DESC").Offset(p.Offset()).Limit(p.PageSize).Find(&logs)
|
||||
|
||||
taskIDList := make([]string, 0)
|
||||
for _, log := range logs {
|
||||
taskIDList = append(taskIDList, log.TaskID)
|
||||
}
|
||||
|
||||
var tasks []models.Task
|
||||
database.DB.Where("id IN ?", taskIDList).Find(&tasks)
|
||||
taskMap := make(map[string]models.Task)
|
||||
for _, t := range tasks {
|
||||
taskMap[t.ID] = t
|
||||
}
|
||||
|
||||
result := make([]vo.TaskLogVO, len(logs))
|
||||
for i, log := range logs {
|
||||
task := taskMap[log.TaskID]
|
||||
taskType := task.Type
|
||||
if taskType == "" {
|
||||
taskType = "task"
|
||||
}
|
||||
result[i] = vo.TaskLogVO{
|
||||
ID: log.ID,
|
||||
TaskID: log.TaskID,
|
||||
TaskName: task.Name,
|
||||
TaskType: taskType,
|
||||
AgentID: log.AgentID,
|
||||
Command: string(log.Command),
|
||||
Status: log.Status,
|
||||
Duration: log.Duration,
|
||||
StartTime: log.StartTime,
|
||||
EndTime: log.EndTime,
|
||||
CreatedAt: log.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
utils.PaginatedResponse(c, result, total, p)
|
||||
}
|
||||
|
||||
// GetLogDetail 获取日志详情
|
||||
// @Summary 获取日志详情
|
||||
// @Description 根据 ID 获取任务日志详细内容(包含输出)
|
||||
// @Tags 日志管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "日志ID"
|
||||
// @Success 200 {object} utils.Response{data=vo.TaskLogVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /logs/{id} [get]
|
||||
func (lc *LogController) GetLogDetail(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的日志ID")
|
||||
return
|
||||
}
|
||||
|
||||
var log models.TaskLog
|
||||
res := database.DB.Where("id = ?", id).Limit(1).Find(&log)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
utils.NotFound(c, "日志不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskLogVO(&log))
|
||||
}
|
||||
|
||||
// ClearLogs 清空日志
|
||||
func (lc *LogController) ClearLogs(c *gin.Context) {
|
||||
var req struct {
|
||||
TaskID *string `json:"task_id"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
query := database.DB.Model(&models.TaskLog{})
|
||||
if req.TaskID != nil && *req.TaskID != "" {
|
||||
query = query.Where("task_id = ?", *req.TaskID)
|
||||
} else {
|
||||
query = query.Where("1 = 1") // Allow delete all without GORM safety block
|
||||
}
|
||||
|
||||
if err := query.Delete(&models.TaskLog{}).Error; err != nil {
|
||||
utils.ServerError(c, "清空日志失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "日志清空成功")
|
||||
}
|
||||
|
||||
// DeleteLog 删除日志
|
||||
func (lc *LogController) DeleteLog(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的日志ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := database.DB.Where("id = ?", id).Delete(&models.TaskLog{}).Error; err != nil {
|
||||
utils.ServerError(c, "删除日志失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "日志已删除")
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type LogSSEController struct{}
|
||||
|
||||
func NewLogSSEController() *LogSSEController {
|
||||
return &LogSSEController{}
|
||||
}
|
||||
|
||||
func (lc *LogSSEController) StreamLog(c *gin.Context) {
|
||||
logIDStr := c.Query("log_id")
|
||||
if logIDStr == "" {
|
||||
c.JSON(400, gin.H{"error": "log_id is required"})
|
||||
return
|
||||
}
|
||||
|
||||
logID := logIDStr
|
||||
|
||||
c.Header("Content-Type", "text/event-stream")
|
||||
c.Header("Cache-Control", "no-cache")
|
||||
c.Header("Connection", "keep-alive")
|
||||
c.Header("Transfer-Encoding", "chunked")
|
||||
// c.Header("Access-Control-Allow-Origin", "*")
|
||||
|
||||
// 1. 检查数据库中是否已结束
|
||||
var taskLog models.TaskLog
|
||||
res := database.DB.Where("id = ?", logID).Limit(1).Find(&taskLog)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
if taskLog.Status != "running" {
|
||||
// 已结束,读取库内日志
|
||||
content, err := utils.DecompressFromBase64(string(taskLog.Output))
|
||||
if err != nil {
|
||||
c.SSEvent("message", gin.H{"text": "解压日志失败: " + err.Error()})
|
||||
c.Writer.Flush()
|
||||
return
|
||||
}
|
||||
c.SSEvent("message", gin.H{"text": content})
|
||||
c.Writer.Flush()
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 未结束或未找到记录,尝试从 TinyLogManager 获取
|
||||
tl := tasks.GetActiveLog(logID)
|
||||
if tl == nil {
|
||||
c.SSEvent("message", gin.H{"text": "未找到正在运行的任务日志"})
|
||||
c.Writer.Flush()
|
||||
return
|
||||
}
|
||||
|
||||
// 发送系统提示
|
||||
c.SSEvent("message", gin.H{"text": fmt.Sprintf("[System] 连接成功,正在监听日志... (LogID: %s)\n", logID)})
|
||||
c.Writer.Flush()
|
||||
|
||||
// 发送最后 100 行
|
||||
lastLines, err := tl.ReadLastLines(100)
|
||||
if err == nil && len(lastLines) > 0 {
|
||||
c.SSEvent("message", gin.H{"text": string(lastLines)})
|
||||
c.Writer.Flush()
|
||||
}
|
||||
|
||||
// 订阅实时更新
|
||||
sub := tl.Subscribe()
|
||||
defer tl.Unsubscribe(sub)
|
||||
|
||||
// 推送更新
|
||||
c.Stream(func(w io.Writer) bool {
|
||||
select {
|
||||
case data, ok := <-sub:
|
||||
if !ok {
|
||||
// 任务结束,尝试刷新最后一次库内完整内容
|
||||
var finalLog models.TaskLog
|
||||
res := database.DB.Where("id = ?", logID).Limit(1).Find(&finalLog)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
content, _ := utils.DecompressFromBase64(string(finalLog.Output))
|
||||
if content != "" {
|
||||
c.SSEvent("message", gin.H{"text": "\n--- 任务已结束 ---\n"})
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
c.SSEvent("message", gin.H{"text": string(data)})
|
||||
return true
|
||||
case <-c.Request.Context().Done():
|
||||
return false
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,165 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type MiseController struct {
|
||||
service *services.MiseService
|
||||
}
|
||||
|
||||
func NewMiseController(service *services.MiseService) *MiseController {
|
||||
return &MiseController{
|
||||
service: service,
|
||||
}
|
||||
}
|
||||
|
||||
// List 获取语言列表
|
||||
func (c *MiseController) List(ctx *gin.Context) {
|
||||
langs, err := c.service.List()
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取语言列表失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, langs)
|
||||
}
|
||||
|
||||
// Sync 同步本地环境到数据库
|
||||
func (c *MiseController) Sync(ctx *gin.Context) {
|
||||
if err := c.service.Sync(); err != nil {
|
||||
utils.ServerError(ctx, "同步本地环境失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
|
||||
// Plugins 获取可用插件列表
|
||||
func (c *MiseController) Plugins(ctx *gin.Context) {
|
||||
plugins, err := c.service.Plugins()
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取插件列表失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, plugins)
|
||||
}
|
||||
|
||||
// Versions 获取指定插件的可用版本列表
|
||||
func (c *MiseController) Versions(ctx *gin.Context) {
|
||||
plugin := ctx.Query("plugin")
|
||||
if plugin == "" {
|
||||
utils.BadRequest(ctx, "参数 plugin 不能为空")
|
||||
return
|
||||
}
|
||||
versions, err := c.service.Versions(plugin)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取版本列表失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, versions)
|
||||
}
|
||||
|
||||
// VerifyCommand 获取验证命令
|
||||
func (c *MiseController) VerifyCommand(ctx *gin.Context) {
|
||||
plugin := ctx.Query("plugin")
|
||||
version := ctx.Query("version")
|
||||
if plugin == "" {
|
||||
utils.BadRequest(ctx, "参数 plugin 不能为空")
|
||||
return
|
||||
}
|
||||
cmd, err := c.service.GetVerifyCommand(plugin, version)
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取验证命令失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, gin.H{"command": cmd})
|
||||
}
|
||||
|
||||
// UseGlobal 设置全局默认版本
|
||||
func (c *MiseController) UseGlobal(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Plugin string `json:"plugin"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
||||
return
|
||||
}
|
||||
if req.Plugin == "" || req.Version == "" {
|
||||
utils.BadRequest(ctx, "参数 plugin 和 version 不能为空")
|
||||
return
|
||||
}
|
||||
if err := c.service.UseGlobal(req.Plugin, req.Version); err != nil {
|
||||
utils.ServerError(ctx, "设置全局版本失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
|
||||
// UnsetGlobal 取消全局默认版本
|
||||
func (c *MiseController) UnsetGlobal(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Plugin string `json:"plugin"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
||||
return
|
||||
}
|
||||
if req.Plugin == "" {
|
||||
utils.BadRequest(ctx, "参数 plugin 不能为空")
|
||||
return
|
||||
}
|
||||
if err := c.service.UnsetGlobal(req.Plugin, req.Version); err != nil {
|
||||
utils.ServerError(ctx, "取消全局版本失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
|
||||
// Envs 获取全局环境变量
|
||||
func (c *MiseController) Envs(ctx *gin.Context) {
|
||||
envs, err := c.service.Envs()
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, "获取全局环境变量失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, envs)
|
||||
}
|
||||
|
||||
// SetEnv 设置全局环境变量
|
||||
func (c *MiseController) SetEnv(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Key string `json:"key"`
|
||||
Value string `json:"value"`
|
||||
}
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
||||
return
|
||||
}
|
||||
if req.Key == "" {
|
||||
utils.BadRequest(ctx, "参数 key 不能为空")
|
||||
return
|
||||
}
|
||||
if err := c.service.SetEnv(req.Key, req.Value); err != nil {
|
||||
utils.ServerError(ctx, "设置环境变量失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
|
||||
// UnsetEnv 取消全局环境变量
|
||||
func (c *MiseController) UnsetEnv(ctx *gin.Context) {
|
||||
key := ctx.Query("key")
|
||||
if key == "" {
|
||||
utils.BadRequest(ctx, "参数 key 不能为空")
|
||||
return
|
||||
}
|
||||
if err := c.service.UnsetEnv(key); err != nil {
|
||||
utils.ServerError(ctx, "取消环境变量失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, nil)
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"runtime"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
type MonitorController struct {
|
||||
executorService *tasks.ExecutorService
|
||||
}
|
||||
|
||||
func NewMonitorController(executorService *tasks.ExecutorService) *MonitorController {
|
||||
return &MonitorController{
|
||||
executorService: executorService,
|
||||
}
|
||||
}
|
||||
|
||||
var monitorUpgrader = websocket.Upgrader{
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
return true // 开发环境允许所有跨域,生产环境可根据配置限制
|
||||
},
|
||||
}
|
||||
|
||||
// GetSystemMonitor 获取系统和内存监控信息 (HTTP)
|
||||
func (mc *MonitorController) GetSystemMonitor(c *gin.Context) {
|
||||
data := mc.getMonitorData()
|
||||
utils.Success(c, data)
|
||||
}
|
||||
|
||||
// MonitorSSE Server-Sent Events 获取系统监控数据
|
||||
func (mc *MonitorController) MonitorSSE(c *gin.Context) {
|
||||
// 设置 SSE 响应头
|
||||
c.Writer.Header().Set("Content-Type", "text/event-stream")
|
||||
c.Writer.Header().Set("Cache-Control", "no-cache")
|
||||
c.Writer.Header().Set("Connection", "keep-alive")
|
||||
c.Writer.Header().Set("Transfer-Encoding", "chunked")
|
||||
|
||||
// 初始发送一次数据
|
||||
if err := mc.sendMonitorDataSSE(c); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(5 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
if err := mc.sendMonitorDataSSE(c); err != nil {
|
||||
return // 客户端断开连接或发送失败
|
||||
}
|
||||
case <-c.Request.Context().Done():
|
||||
return // 连接已断开,立即退出
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (mc *MonitorController) sendMonitorDataSSE(c *gin.Context) error {
|
||||
data := mc.getMonitorData()
|
||||
// 使用 Gin 提供的 SSE 方法
|
||||
c.SSEvent("message", gin.H{
|
||||
"code": 200,
|
||||
"data": data,
|
||||
"msg": "success",
|
||||
})
|
||||
c.Writer.Flush()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (mc *MonitorController) getMonitorData() gin.H {
|
||||
rt := services.GetMonitorService().GetRuntimeMetrics()
|
||||
m := rt.MemStats
|
||||
|
||||
// 调用统一的监控服务获取物理机指标
|
||||
metrics := services.GetMonitorService().GetHostMetrics()
|
||||
|
||||
return gin.H{
|
||||
"env": gin.H{
|
||||
"os": runtime.GOOS,
|
||||
"arch": runtime.GOARCH,
|
||||
"go_version": runtime.Version(),
|
||||
"num_cpu": runtime.NumCPU(),
|
||||
"goroutines": rt.NumGoroutine,
|
||||
},
|
||||
"host": gin.H{
|
||||
"cpu_percent": metrics.CPUPercent,
|
||||
"mem_total": metrics.VMem.Total,
|
||||
"mem_used": metrics.VMem.Used,
|
||||
"mem_percent": metrics.VMem.UsedPercent,
|
||||
"disk_total": metrics.DiskUsage.Total,
|
||||
"disk_used": metrics.DiskUsage.Used,
|
||||
"disk_percent": metrics.DiskUsage.UsedPercent,
|
||||
"uptime": metrics.HostInfo.Uptime,
|
||||
"platform": metrics.HostInfo.Platform + " " + metrics.HostInfo.PlatformVersion,
|
||||
},
|
||||
"mem": gin.H{
|
||||
"alloc": m.Alloc,
|
||||
"total_alloc": m.TotalAlloc,
|
||||
"sys": m.Sys,
|
||||
"lookups": m.Lookups,
|
||||
"mallocs": m.Mallocs,
|
||||
"frees": m.Frees,
|
||||
},
|
||||
"heap": gin.H{
|
||||
"heap_alloc": m.HeapAlloc,
|
||||
"heap_sys": m.HeapSys,
|
||||
"heap_idle": m.HeapIdle,
|
||||
"heap_inuse": m.HeapInuse,
|
||||
"heap_released": m.HeapReleased,
|
||||
"heap_objects": m.HeapObjects,
|
||||
},
|
||||
"gc": gin.H{
|
||||
"next_gc": m.NextGC,
|
||||
"last_gc": m.LastGC,
|
||||
"pause_total_ns": m.PauseTotalNs,
|
||||
"num_gc": m.NumGC,
|
||||
},
|
||||
"scheduler": gin.H{
|
||||
"scheduled": mc.executorService.GetScheduledCount(),
|
||||
"running": mc.executorService.GetRunningCount(),
|
||||
"queue_size": mc.executorService.GetScheduler().GetQueueSize(),
|
||||
"worker_count": mc.executorService.GetScheduler().GetConfig().WorkerCount,
|
||||
"workers": mc.executorService.GetScheduler().GetWorkerStatuses(),
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type NotificationController struct {
|
||||
notifyService *services.NotificationService
|
||||
}
|
||||
|
||||
func NewNotificationController() *NotificationController {
|
||||
return &NotificationController{
|
||||
notifyService: services.NewNotificationService(),
|
||||
}
|
||||
}
|
||||
|
||||
// GetChannelTypes 获取支持的渠道类型
|
||||
func (nc *NotificationController) GetChannelTypes(c *gin.Context) {
|
||||
utils.Success(c, gin.H{
|
||||
"channel_types": services.SupportedChannelTypes,
|
||||
"event_types": services.SupportedEvents,
|
||||
})
|
||||
}
|
||||
|
||||
// GetChannels 获取所有渠道
|
||||
func (nc *NotificationController) GetChannels(c *gin.Context) {
|
||||
channels := nc.notifyService.GetChannels()
|
||||
utils.Success(c, channels)
|
||||
}
|
||||
|
||||
// SaveChannel 保存/更新渠道
|
||||
func (nc *NotificationController) SaveChannel(c *gin.Context) {
|
||||
var req services.NotifyChannel
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" || req.Type == "" {
|
||||
utils.BadRequest(c, "渠道名称和类型不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
if err := nc.notifyService.SaveChannel(req); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// DeleteChannel 删除渠道
|
||||
func (nc *NotificationController) DeleteChannel(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "缺少渠道ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := nc.notifyService.DeleteChannel(id); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
// TestChannel 测试渠道
|
||||
func (nc *NotificationController) TestChannel(c *gin.Context) {
|
||||
var req services.NotifyChannel
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
result := nc.notifyService.SendToChannel(req, &services.NotifyMessage{
|
||||
Title: "🔔 任务池测试通知",
|
||||
Text: "如果你看到这条消息,说明通知渠道配置正确!",
|
||||
})
|
||||
|
||||
utils.Success(c, result)
|
||||
}
|
||||
|
||||
// GetBindings 获取事件绑定列表
|
||||
func (nc *NotificationController) GetBindings(c *gin.Context) {
|
||||
bindings := nc.notifyService.GetBindings()
|
||||
utils.Success(c, bindings)
|
||||
}
|
||||
|
||||
// SaveBinding 保存事件绑定
|
||||
func (nc *NotificationController) SaveBinding(c *gin.Context) {
|
||||
var req struct {
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"`
|
||||
Event string `json:"event"`
|
||||
WayID string `json:"way_id"`
|
||||
DataID string `json:"data_id"`
|
||||
Extra models.BigText `json:"extra"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Type == "" || req.Event == "" || req.WayID == "" {
|
||||
utils.BadRequest(c, "类型、事件和渠道ID不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
binding := &models.NotifyBinding{
|
||||
ID: req.ID,
|
||||
Type: req.Type,
|
||||
Event: req.Event,
|
||||
WayID: req.WayID,
|
||||
DataID: req.DataID,
|
||||
Extra: req.Extra,
|
||||
}
|
||||
|
||||
if err := nc.notifyService.SaveBinding(binding); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, binding)
|
||||
}
|
||||
|
||||
// DeleteBinding 删除事件绑定
|
||||
func (nc *NotificationController) DeleteBinding(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "缺少绑定ID")
|
||||
return
|
||||
}
|
||||
|
||||
if err := nc.notifyService.DeleteBinding(id); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
// BatchSaveBindings 批量保存事件绑定
|
||||
func (nc *NotificationController) BatchSaveBindings(c *gin.Context) {
|
||||
var req struct {
|
||||
Type string `json:"type"`
|
||||
DataID string `json:"data_id"`
|
||||
Bindings []models.NotifyBinding `json:"bindings"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Type == "" {
|
||||
utils.BadRequest(c, "类型不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
if err := nc.notifyService.BatchSaveBindings(req.Type, req.DataID, req.Bindings); err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// SendNotification API 发送通知(供脚本调用)
|
||||
func (nc *NotificationController) SendNotification(c *gin.Context) {
|
||||
var req struct {
|
||||
ChannelID string `json:"channel_id"`
|
||||
Title string `json:"title"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if req.ChannelID == "" || req.Title == "" {
|
||||
utils.BadRequest(c, "channel_id 和 title 不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
result := nc.notifyService.SendByChannelID(req.ChannelID, &services.NotifyMessage{
|
||||
Title: req.Title,
|
||||
Text: req.Text,
|
||||
})
|
||||
|
||||
utils.Success(c, result)
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type ScriptController struct {
|
||||
scriptService *services.ScriptService
|
||||
}
|
||||
|
||||
func NewScriptController(scriptService *services.ScriptService) *ScriptController {
|
||||
return &ScriptController{scriptService: scriptService}
|
||||
}
|
||||
|
||||
// CreateScript 创建脚本
|
||||
// @Summary 创建脚本
|
||||
// @Description 创建一个新的脚本
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param body body object true "脚本信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.ScriptVO}
|
||||
// @Router /scripts [post]
|
||||
func (sc *ScriptController) CreateScript(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
Content string `json:"content" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
script := sc.scriptService.CreateScript(req.Name, req.Content, userID)
|
||||
utils.Success(c, vo.ToScriptVO(script))
|
||||
}
|
||||
|
||||
// GetScripts 获取脚本列表
|
||||
// @Summary 获取脚本列表
|
||||
// @Description 获取当前用户的所有脚本(内容字段为空)
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Success 200 {object} utils.Response{data=[]vo.ScriptVO}
|
||||
// @Router /scripts [get]
|
||||
func (sc *ScriptController) GetScripts(c *gin.Context) {
|
||||
userID := c.GetString("userID")
|
||||
scripts := sc.scriptService.GetScriptsByUserID(userID)
|
||||
vos := vo.ToScriptVOListFromModels(scripts)
|
||||
for i := range vos {
|
||||
vos[i].Content = "" // 列表不返回内容
|
||||
}
|
||||
utils.Success(c, vos)
|
||||
}
|
||||
|
||||
// GetScript 获取脚本详情
|
||||
// @Summary 获取脚本详情
|
||||
// @Description 根据 ID 获取脚本详情
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "脚本ID"
|
||||
// @Success 200 {object} utils.Response{data=vo.ScriptVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /scripts/{id} [get]
|
||||
func (sc *ScriptController) GetScript(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的脚本ID")
|
||||
return
|
||||
}
|
||||
|
||||
script := sc.scriptService.GetScriptByID(id)
|
||||
if script == nil {
|
||||
utils.NotFound(c, "脚本不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToScriptVO(script))
|
||||
}
|
||||
|
||||
// UpdateScript 更新脚本
|
||||
// @Summary 更新脚本
|
||||
// @Description 根据 ID 更新脚本信息
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "脚本ID"
|
||||
// @Param body body object true "脚本更新信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.ScriptVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /scripts/{id} [put]
|
||||
func (sc *ScriptController) UpdateScript(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的脚本ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
script := sc.scriptService.UpdateScript(id, req.Name, req.Content)
|
||||
if script == nil {
|
||||
utils.NotFound(c, "脚本不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToScriptVO(script))
|
||||
}
|
||||
|
||||
// DeleteScript 删除脚本
|
||||
// @Summary 删除脚本
|
||||
// @Description 根据 ID 删除脚本
|
||||
// @Tags 脚本管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "脚本ID"
|
||||
// @Success 200 {object} utils.Response
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /scripts/{id} [delete]
|
||||
func (sc *ScriptController) DeleteScript(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的脚本ID")
|
||||
return
|
||||
}
|
||||
|
||||
success := sc.scriptService.DeleteScript(id)
|
||||
if !success {
|
||||
utils.NotFound(c, "脚本不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
@@ -0,0 +1,624 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/eventbus"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/tunnel"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/shirou/gopsutil/v3/process"
|
||||
)
|
||||
|
||||
type SettingsController struct {
|
||||
userService *services.UserService
|
||||
settingsService *services.SettingsService
|
||||
loginLogService *services.LoginLogService
|
||||
backupService *services.BackupService
|
||||
executorService *tasks.ExecutorService
|
||||
}
|
||||
|
||||
func NewSettingsController(userService *services.UserService, loginLogService *services.LoginLogService, executorService *tasks.ExecutorService) *SettingsController {
|
||||
return &SettingsController{
|
||||
userService: userService,
|
||||
settingsService: services.NewSettingsService(),
|
||||
loginLogService: loginLogService,
|
||||
backupService: services.NewBackupService(),
|
||||
executorService: executorService,
|
||||
}
|
||||
}
|
||||
|
||||
// ChangePassword 修改密码及账号信息
|
||||
func (sc *SettingsController) ChangePassword(c *gin.Context) {
|
||||
// 演示模式下禁止修改
|
||||
if constant.DemoMode {
|
||||
utils.BadRequest(c, "演示模式下不能修改账号或密码")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
OldUsername string `json:"old_username"`
|
||||
Username string `json:"username"`
|
||||
OldPassword string `json:"old_password" binding:"required"`
|
||||
NewPassword string `json:"new_password"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
userID := c.GetString("userID")
|
||||
var user *models.User
|
||||
res := database.DB.Where("id = ?", userID).Limit(1).Find(&user)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
utils.NotFound(c, "用户不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 统一校验原账密
|
||||
if req.OldUsername != "" && req.OldUsername != user.Username {
|
||||
utils.BadRequest(c, "原账号不正确")
|
||||
return
|
||||
}
|
||||
|
||||
if !sc.userService.AuthenticateUser(user.Username, req.OldPassword) {
|
||||
utils.BadRequest(c, "原密码错误")
|
||||
return
|
||||
}
|
||||
|
||||
var updated bool
|
||||
var logoutRequired bool
|
||||
|
||||
// 1. 处理用户名修改
|
||||
if req.Username != "" && req.Username != user.Username {
|
||||
if err := sc.userService.UpdateAccount(user.ID, req.Username); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
updated = true
|
||||
logoutRequired = true
|
||||
}
|
||||
|
||||
// 2. 处理密码修改
|
||||
if req.NewPassword != "" {
|
||||
if len(req.NewPassword) < 6 {
|
||||
utils.BadRequest(c, "新密码至少6位")
|
||||
return
|
||||
}
|
||||
if err := sc.userService.UpdatePassword(user.ID, req.NewPassword); err != nil {
|
||||
utils.ServerError(c, "修改密码失败")
|
||||
return
|
||||
}
|
||||
updated = true
|
||||
logoutRequired = true
|
||||
}
|
||||
|
||||
if !updated {
|
||||
utils.SuccessMsg(c, "未检测到变更内容")
|
||||
return
|
||||
}
|
||||
|
||||
eventbus.DefaultBus.Publish(eventbus.Event{
|
||||
Type: constant.EventPasswordChanged,
|
||||
Payload: map[string]interface{}{
|
||||
"username": user.Username,
|
||||
},
|
||||
})
|
||||
|
||||
msg := "保存成功"
|
||||
if logoutRequired {
|
||||
msg += ",请重新登录"
|
||||
}
|
||||
utils.SuccessMsg(c, msg)
|
||||
}
|
||||
|
||||
// CleanLogs 清理日志 - 已移除,改为任务级别的日志清理配置
|
||||
|
||||
// GetSiteSettings 获取站点设置
|
||||
func (sc *SettingsController) GetSiteSettings(c *gin.Context) {
|
||||
settings := sc.settingsService.GetSection(constant.SectionSite)
|
||||
|
||||
// 纠正数据库中的空值,防止因配置冲突被意外置空
|
||||
if settings[constant.KeyTitle] == "" {
|
||||
settings[constant.KeyTitle] = "任务池"
|
||||
sc.settingsService.Set(constant.SectionSite, constant.KeyTitle, "任务池")
|
||||
}
|
||||
if settings[constant.KeySubtitle] == "" {
|
||||
settings[constant.KeySubtitle] = "极致轻量、高性能的自动化任务调度平台"
|
||||
sc.settingsService.Set(constant.SectionSite, constant.KeySubtitle, "极致轻量、高性能的自动化任务调度平台")
|
||||
}
|
||||
if settings[constant.KeyIcon] == "" {
|
||||
settings[constant.KeyIcon] = constant.DefaultIcon
|
||||
sc.settingsService.Set(constant.SectionSite, constant.KeyIcon, constant.DefaultIcon)
|
||||
}
|
||||
|
||||
// 解析 JSON 格式的 OpenAPI Token
|
||||
if tokenJson, ok := settings[constant.KeyOpenapiToken]; ok && tokenJson != "" {
|
||||
var tokenConfig vo.TokenConfig
|
||||
if err := json.Unmarshal([]byte(tokenJson), &tokenConfig); err == nil {
|
||||
settings["openapi_token"] = tokenConfig.Token
|
||||
settings["openapi_token_expire"] = tokenConfig.ExpireAt
|
||||
if tokenConfig.Enabled {
|
||||
settings["openapi_enabled"] = "true"
|
||||
} else {
|
||||
settings["openapi_enabled"] = "false"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 获取日志清理配置
|
||||
settings["system_notice_days"] = sc.settingsService.Get(constant.SectionSystem, constant.KeySystemNoticeDays)
|
||||
settings["system_notice_max_count"] = sc.settingsService.Get(constant.SectionSystem, constant.KeySystemNoticeMaxCount)
|
||||
settings["push_log_days"] = sc.settingsService.Get(constant.SectionSystem, constant.KeyPushLogDays)
|
||||
settings["push_log_max_count"] = sc.settingsService.Get(constant.SectionSystem, constant.KeyPushLogMaxCount)
|
||||
settings["login_log_days"] = sc.settingsService.Get(constant.SectionSystem, constant.KeyLoginLogDays)
|
||||
settings["login_log_max_count"] = sc.settingsService.Get(constant.SectionSystem, constant.KeyLoginLogMaxCount)
|
||||
settings["scheduler_log_days"] = sc.settingsService.Get(constant.SectionSystem, constant.KeySchedulerLogDays)
|
||||
settings["scheduler_log_max_count"] = sc.settingsService.Get(constant.SectionSystem, constant.KeySchedulerLogMaxCount)
|
||||
|
||||
utils.Success(c, settings)
|
||||
}
|
||||
|
||||
// GetPublicSiteSettings 获取公开的站点设置(无需认证)
|
||||
func (sc *SettingsController) GetPublicSiteSettings(c *gin.Context) {
|
||||
settings := sc.settingsService.GetSection(constant.SectionSite)
|
||||
|
||||
title := settings[constant.KeyTitle]
|
||||
if title == "" {
|
||||
title = "任务池"
|
||||
}
|
||||
subtitle := settings[constant.KeySubtitle]
|
||||
if subtitle == "" {
|
||||
subtitle = "极致轻量、高性能的自动化任务调度平台"
|
||||
}
|
||||
icon := settings[constant.KeyIcon]
|
||||
if icon == "" {
|
||||
icon = constant.DefaultIcon
|
||||
}
|
||||
|
||||
// 只返回公开信息
|
||||
utils.Success(c, gin.H{
|
||||
constant.KeyTitle: title,
|
||||
constant.KeySubtitle: subtitle,
|
||||
constant.KeyIcon: icon,
|
||||
"demo_mode": constant.DemoMode,
|
||||
})
|
||||
}
|
||||
|
||||
// UpdateSiteSettings 更新站点设置
|
||||
func (sc *SettingsController) UpdateSiteSettings(c *gin.Context) {
|
||||
var req struct {
|
||||
Title string `json:"title"`
|
||||
Subtitle string `json:"subtitle"`
|
||||
Icon string `json:"icon"`
|
||||
PageSize string `json:"page_size"`
|
||||
CookieDays string `json:"cookie_days"`
|
||||
OpenapiEnabled bool `json:"openapi_enabled"`
|
||||
OpenapiToken string `json:"openapi_token"`
|
||||
OpenapiTokenExpire string `json:"openapi_token_expire"`
|
||||
SystemNoticeDays string `json:"system_notice_days"`
|
||||
SystemNoticeMaxCount string `json:"system_notice_max_count"`
|
||||
PushLogDays string `json:"push_log_days"`
|
||||
PushLogMaxCount string `json:"push_log_max_count"`
|
||||
LoginLogDays string `json:"login_log_days"`
|
||||
LoginLogMaxCount string `json:"login_log_max_count"`
|
||||
SchedulerLogDays string `json:"scheduler_log_days"`
|
||||
SchedulerLogMaxCount string `json:"scheduler_log_max_count"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
openapiTokenJson := ""
|
||||
if req.OpenapiToken != "" || req.OpenapiTokenExpire != "" || req.OpenapiEnabled {
|
||||
tokenConfig := vo.TokenConfig{
|
||||
Enabled: req.OpenapiEnabled,
|
||||
Token: req.OpenapiToken,
|
||||
ExpireAt: req.OpenapiTokenExpire,
|
||||
}
|
||||
if b, err := json.Marshal(tokenConfig); err == nil {
|
||||
openapiTokenJson = string(b)
|
||||
}
|
||||
}
|
||||
|
||||
values := map[string]string{
|
||||
constant.KeyTitle: req.Title,
|
||||
constant.KeySubtitle: req.Subtitle,
|
||||
constant.KeyIcon: req.Icon,
|
||||
constant.KeyPageSize: req.PageSize,
|
||||
constant.KeyCookieDays: req.CookieDays,
|
||||
constant.KeyOpenapiToken: openapiTokenJson,
|
||||
}
|
||||
|
||||
if err := sc.settingsService.SetSection(constant.SectionSite, values); err != nil {
|
||||
utils.ServerError(c, "保存失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 保存日志清理配置
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeySystemNoticeDays, req.SystemNoticeDays)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeySystemNoticeMaxCount, req.SystemNoticeMaxCount)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeyPushLogDays, req.PushLogDays)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeyPushLogMaxCount, req.PushLogMaxCount)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeyLoginLogDays, req.LoginLogDays)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeyLoginLogMaxCount, req.LoginLogMaxCount)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeySchedulerLogDays, req.SchedulerLogDays)
|
||||
sc.settingsService.Set(constant.SectionSystem, constant.KeySchedulerLogMaxCount, req.SchedulerLogMaxCount)
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// GenerateOpenapiToken 随机生成OpenAPI Token
|
||||
func (sc *SettingsController) GenerateOpenapiToken(c *gin.Context) {
|
||||
utils.Success(c, gin.H{
|
||||
"token": strings.ToLower(utils.RandomString(32)),
|
||||
})
|
||||
}
|
||||
|
||||
// GetSchedulerSettings 获取调度设置
|
||||
func (sc *SettingsController) GetSchedulerSettings(c *gin.Context) {
|
||||
settings := sc.settingsService.GetSection(constant.SectionScheduler)
|
||||
utils.Success(c, settings)
|
||||
}
|
||||
|
||||
// UpdateSchedulerSettings 更新调度设置
|
||||
func (sc *SettingsController) UpdateSchedulerSettings(c *gin.Context) {
|
||||
var req struct {
|
||||
WorkerCount string `json:"worker_count"`
|
||||
QueueSize string `json:"queue_size"`
|
||||
RateInterval string `json:"rate_interval"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
var workerCount int
|
||||
if _, err := fmt.Sscanf(req.WorkerCount, "%d", &workerCount); err != nil || workerCount < 1 || workerCount > 1000 {
|
||||
utils.BadRequest(c, "工作线程数必须在 1 至 1000 之间")
|
||||
return
|
||||
}
|
||||
|
||||
var queueSize int
|
||||
if _, err := fmt.Sscanf(req.QueueSize, "%d", &queueSize); err != nil || queueSize < 1 || queueSize > 50000 {
|
||||
utils.BadRequest(c, "等待队列容量必须在 1 至 50000 之间")
|
||||
return
|
||||
}
|
||||
|
||||
var rateInterval int
|
||||
if _, err := fmt.Sscanf(req.RateInterval, "%d", &rateInterval); err != nil || rateInterval < 1 {
|
||||
utils.BadRequest(c, "限频间隔必须为正整数")
|
||||
return
|
||||
}
|
||||
|
||||
values := map[string]string{
|
||||
constant.KeyWorkerCount: req.WorkerCount,
|
||||
constant.KeyQueueSize: req.QueueSize,
|
||||
constant.KeyRateInterval: req.RateInterval,
|
||||
}
|
||||
|
||||
if err := sc.settingsService.SetSection(constant.SectionScheduler, values); err != nil {
|
||||
utils.ServerError(c, "保存失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 重新加载 executor service
|
||||
if sc.executorService != nil {
|
||||
sc.executorService.Reload()
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// GetPaths 获取系统路径信息
|
||||
func (sc *SettingsController) GetPaths(c *gin.Context) {
|
||||
absScriptsDir, _ := filepath.Abs(constant.ScriptsWorkDir)
|
||||
utils.Success(c, gin.H{
|
||||
"scripts_dir": absScriptsDir,
|
||||
})
|
||||
}
|
||||
|
||||
// GetAbout 获取关于信息
|
||||
func (sc *SettingsController) GetAbout(c *gin.Context) {
|
||||
var taskCount, logCount, envCount int64
|
||||
database.DB.Model(&models.Task{}).Count(&taskCount)
|
||||
database.DB.Model(&models.TaskLog{}).Count(&logCount)
|
||||
database.DB.Model(&models.EnvironmentVariable{}).Count(&envCount)
|
||||
|
||||
// 内存使用
|
||||
memUsage := "N/A"
|
||||
if p, err := process.NewProcess(int32(os.Getpid())); err == nil {
|
||||
if memInfo, err := p.MemoryInfo(); err == nil {
|
||||
memUsage = formatBytes(memInfo.RSS)
|
||||
}
|
||||
}
|
||||
|
||||
// 运行时间
|
||||
uptime := formatDuration(time.Since(constant.StartTime))
|
||||
|
||||
// 获取远程最新版本
|
||||
remoteVersion := ""
|
||||
client := &http.Client{Timeout: 2 * time.Second}
|
||||
req, err := http.NewRequest("GET", "https://api.github.com/repos/engigu/taskpool/releases/latest", nil)
|
||||
if err == nil {
|
||||
req.Header.Set("User-Agent", "taskpool")
|
||||
if resp, err := client.Do(req); err == nil {
|
||||
defer resp.Body.Close()
|
||||
var release struct {
|
||||
TagName string `json:"tag_name"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&release); err == nil {
|
||||
remoteVersion = release.TagName
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"version": constant.Version,
|
||||
"remote_version": remoteVersion,
|
||||
"build_time": constant.BuildTime,
|
||||
"mem_usage": memUsage,
|
||||
"goroutines": runtime.NumGoroutine(),
|
||||
"uptime": uptime,
|
||||
"task_count": taskCount,
|
||||
"log_count": logCount,
|
||||
"env_count": envCount,
|
||||
})
|
||||
}
|
||||
|
||||
// GetChangelog 获取更新日志
|
||||
func (sc *SettingsController) GetChangelog(c *gin.Context) {
|
||||
content, err := os.ReadFile("docs/guide/changelog.md")
|
||||
if err != nil {
|
||||
utils.Success(c, "暂无更新日志")
|
||||
return
|
||||
}
|
||||
utils.Success(c, string(content))
|
||||
}
|
||||
|
||||
// formatBytes 格式化字节数
|
||||
func formatBytes(bytes uint64) string {
|
||||
const unit = 1024
|
||||
if bytes < unit {
|
||||
return fmt.Sprintf("%d B", bytes)
|
||||
}
|
||||
div, exp := uint64(unit), 0
|
||||
for n := bytes / unit; n >= unit; n /= unit {
|
||||
div *= unit
|
||||
exp++
|
||||
}
|
||||
return fmt.Sprintf("%.1f %cB", float64(bytes)/float64(div), "KMGTPE"[exp])
|
||||
}
|
||||
|
||||
// formatDuration 格式化时间间隔
|
||||
func formatDuration(d time.Duration) string {
|
||||
days := int(d.Hours()) / 24
|
||||
hours := int(d.Hours()) % 24
|
||||
minutes := int(d.Minutes()) % 60
|
||||
seconds := int(d.Seconds()) % 60
|
||||
|
||||
if days > 0 {
|
||||
return fmt.Sprintf("%d天%d小时%d分钟%d秒", days, hours, minutes, seconds)
|
||||
}
|
||||
if hours > 0 {
|
||||
return fmt.Sprintf("%d小时%d分钟%d秒", hours, minutes, seconds)
|
||||
}
|
||||
if minutes > 0 {
|
||||
return fmt.Sprintf("%d分钟%d秒", minutes, seconds)
|
||||
}
|
||||
return fmt.Sprintf("%d秒", seconds)
|
||||
}
|
||||
|
||||
// GetLoginLogs 获取登录日志
|
||||
func (sc *SettingsController) GetLoginLogs(c *gin.Context) {
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "10"))
|
||||
username := c.Query("username")
|
||||
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 || pageSize > 100 {
|
||||
pageSize = 10
|
||||
}
|
||||
|
||||
logs, total, err := sc.loginLogService.List(page, pageSize, username)
|
||||
if err != nil {
|
||||
utils.ServerError(c, "获取登录日志失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 将 AppLog 转换为 LoginLogVO 返回,保持前端兼容性
|
||||
vos := make([]*vo.LoginLogVO, len(logs))
|
||||
for i, log := range logs {
|
||||
vos[i] = &vo.LoginLogVO{
|
||||
ID: log.ID,
|
||||
Username: log.Title,
|
||||
IP: log.RefID,
|
||||
UserAgent: string(log.Content),
|
||||
Status: log.Status,
|
||||
Message: string(log.ErrorMsg),
|
||||
CreatedAt: log.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, utils.PaginationData{
|
||||
Data: vos,
|
||||
Total: total,
|
||||
Page: page,
|
||||
PageSize: pageSize,
|
||||
})
|
||||
}
|
||||
|
||||
// CreateBackup 创建备份
|
||||
func (sc *SettingsController) CreateBackup(c *gin.Context) {
|
||||
_, err := sc.backupService.CreateBackup()
|
||||
if err != nil {
|
||||
utils.ServerError(c, "创建备份失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
utils.SuccessMsg(c, "备份创建成功")
|
||||
}
|
||||
|
||||
// GetBackupStatus 获取备份状态
|
||||
func (sc *SettingsController) GetBackupStatus(c *gin.Context) {
|
||||
filePath := sc.backupService.GetBackupFile()
|
||||
var backupTime string
|
||||
if filePath != "" {
|
||||
if info, err := os.Stat(filePath); err == nil {
|
||||
backupTime = info.ModTime().Format("2006-01-02 15:04:05")
|
||||
}
|
||||
}
|
||||
utils.Success(c, gin.H{
|
||||
"has_backup": filePath != "",
|
||||
"backup_time": backupTime,
|
||||
})
|
||||
}
|
||||
|
||||
// DownloadBackup 下载备份文件
|
||||
func (sc *SettingsController) DownloadBackup(c *gin.Context) {
|
||||
filePath := sc.backupService.GetBackupFile()
|
||||
if filePath == "" {
|
||||
utils.NotFound(c, "没有可下载的备份")
|
||||
return
|
||||
}
|
||||
|
||||
// 检查文件是否存在
|
||||
if _, err := os.Stat(filePath); os.IsNotExist(err) {
|
||||
sc.backupService.ClearBackup()
|
||||
utils.NotFound(c, "备份文件不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 设置响应头
|
||||
c.Header("Content-Disposition", "attachment; filename="+filepath.Base(filePath))
|
||||
c.Header("Content-Type", "application/zip")
|
||||
c.File(filePath)
|
||||
|
||||
// 下载后清除备份记录和文件
|
||||
go func() {
|
||||
time.Sleep(time.Minute * 5) // 等待下载完成
|
||||
sc.backupService.ClearBackup()
|
||||
}()
|
||||
}
|
||||
|
||||
// RestoreBackup 恢复备份
|
||||
func (sc *SettingsController) RestoreBackup(c *gin.Context) {
|
||||
file, err := c.FormFile("file")
|
||||
if err != nil {
|
||||
utils.BadRequest(c, "请上传备份文件")
|
||||
return
|
||||
}
|
||||
|
||||
// 保存上传的文件
|
||||
tempPath := filepath.Join(os.TempDir(), file.Filename)
|
||||
if err := c.SaveUploadedFile(file, tempPath); err != nil {
|
||||
utils.ServerError(c, "保存文件失败")
|
||||
return
|
||||
}
|
||||
defer os.Remove(tempPath)
|
||||
|
||||
// 恢复备份
|
||||
if err := sc.backupService.Restore(tempPath); err != nil {
|
||||
utils.ServerError(c, "恢复失败: "+err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "恢复成功")
|
||||
}
|
||||
|
||||
// GetSectionSettings 获取指定 section 的所有设置
|
||||
func (sc *SettingsController) GetSectionSettings(c *gin.Context) {
|
||||
section := c.Param("section")
|
||||
if section == "" {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
settings := sc.settingsService.GetSection(section)
|
||||
utils.Success(c, settings)
|
||||
}
|
||||
|
||||
// UpdateSectionSettings 批量更新指定 section 的设置
|
||||
func (sc *SettingsController) UpdateSectionSettings(c *gin.Context) {
|
||||
section := c.Param("section")
|
||||
if section == "" {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
var values map[string]string
|
||||
if err := c.ShouldBindJSON(&values); err != nil {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
if err := sc.settingsService.SetSection(section, values); err != nil {
|
||||
utils.ServerError(c, "更新失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 当互联配置发生改变时,通知 tunnel 模块立刻应用新角色,启动或停止相关的后台协程
|
||||
if section == constant.SectionInterconnect {
|
||||
if role, ok := values[constant.KeyInterconnectRole]; ok {
|
||||
tunnel.ApplyRole(role)
|
||||
}
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "保存成功")
|
||||
}
|
||||
|
||||
// GetSetting 获取单个设置值
|
||||
func (sc *SettingsController) GetSetting(c *gin.Context) {
|
||||
section := c.Param("section")
|
||||
key := c.Param("key")
|
||||
|
||||
if section == "" || key == "" {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
value := sc.settingsService.Get(section, key)
|
||||
utils.Success(c, value)
|
||||
}
|
||||
|
||||
// GenerateSettingToken 为指定设置生成随机token
|
||||
func (sc *SettingsController) GenerateSettingToken(c *gin.Context) {
|
||||
section := c.Param("section")
|
||||
key := c.Param("key")
|
||||
|
||||
if section == "" || key == "" {
|
||||
utils.BadRequest(c, "参数错误")
|
||||
return
|
||||
}
|
||||
|
||||
// 生成32位随机token
|
||||
token := strings.ToLower(utils.RandomString(32))
|
||||
|
||||
// 保存到数据库
|
||||
if err := sc.settingsService.Set(section, key, token); err != nil {
|
||||
utils.ServerError(c, "保存失败")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, token)
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
type SystemWSController struct {
|
||||
manager *services.SystemWSManager
|
||||
}
|
||||
|
||||
func NewSystemWSController() *SystemWSController {
|
||||
return &SystemWSController{
|
||||
manager: services.GetSystemWSManager(),
|
||||
}
|
||||
}
|
||||
|
||||
func (sc *SystemWSController) HandleEvents(c *gin.Context) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
logger.Errorf("[SystemWS] 升级 WebSocket 失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
client := sc.manager.Register(conn)
|
||||
defer sc.manager.Unregister(client)
|
||||
|
||||
// 启动写循环
|
||||
go sc.writeLoop(client)
|
||||
|
||||
// 启动读循环 (主要用于检测连接断开和维持心跳)
|
||||
sc.readLoop(client)
|
||||
}
|
||||
|
||||
func (sc *SystemWSController) readLoop(client *services.ClientConnection) {
|
||||
defer client.Close()
|
||||
|
||||
client.Conn.SetReadLimit(constant.MaxMessageSize)
|
||||
client.Conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
client.Conn.SetPongHandler(func(string) error {
|
||||
client.Conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
return nil
|
||||
})
|
||||
|
||||
for {
|
||||
_, _, err := client.Conn.ReadMessage()
|
||||
if err != nil {
|
||||
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) {
|
||||
logger.Warnf("[SystemWS] 客户端异常断开: %v", err)
|
||||
}
|
||||
break
|
||||
}
|
||||
// 暂时不处理来自前端的消息,前端仅作为接收方
|
||||
}
|
||||
}
|
||||
|
||||
func (sc *SystemWSController) writeLoop(client *services.ClientConnection) {
|
||||
ticker := time.NewTicker(constant.PingPeriod)
|
||||
defer func() {
|
||||
ticker.Stop()
|
||||
client.Close()
|
||||
}()
|
||||
|
||||
for {
|
||||
select {
|
||||
case message, ok := <-client.Send:
|
||||
client.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
if !ok {
|
||||
// 通道关闭
|
||||
client.Conn.WriteMessage(websocket.CloseMessage, []byte{})
|
||||
return
|
||||
}
|
||||
|
||||
if err := client.Conn.WriteMessage(websocket.TextMessage, message); err != nil {
|
||||
return
|
||||
}
|
||||
case <-ticker.C:
|
||||
client.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
if err := client.Conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,872 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/database"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/models/vo"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"os"
|
||||
)
|
||||
|
||||
type TaskController struct {
|
||||
taskService *tasks.TaskService
|
||||
executorService *tasks.ExecutorService
|
||||
agentWSManager *services.AgentWSManager
|
||||
}
|
||||
|
||||
func NewTaskController(taskService *tasks.TaskService, executorService *tasks.ExecutorService) *TaskController {
|
||||
return &TaskController{
|
||||
taskService: taskService,
|
||||
executorService: executorService,
|
||||
agentWSManager: services.GetAgentWSManager(),
|
||||
}
|
||||
}
|
||||
|
||||
// resolveWorkDir 将相对路径转换为绝对路径
|
||||
func resolveWorkDir(workDir string) string {
|
||||
if workDir == "" {
|
||||
// 空则使用默认 scripts 目录
|
||||
absPath, err := filepath.Abs(constant.ScriptsWorkDir)
|
||||
if err != nil {
|
||||
return constant.ScriptsWorkDir
|
||||
}
|
||||
return absPath
|
||||
}
|
||||
// 如果已经是绝对路径,直接返回
|
||||
if strings.HasPrefix(workDir, constant.ScriptsDirPlaceholder) {
|
||||
return workDir
|
||||
}
|
||||
if filepath.IsAbs(workDir) {
|
||||
return workDir
|
||||
}
|
||||
// 相对路径,基于 scripts 目录
|
||||
fullPath := filepath.Join(constant.ScriptsWorkDir, workDir)
|
||||
absPath, err := filepath.Abs(fullPath)
|
||||
if err != nil {
|
||||
return fullPath
|
||||
}
|
||||
return absPath
|
||||
}
|
||||
|
||||
// isValidDirName 校验目录名是否合法
|
||||
func isValidDirName(dirName string) bool {
|
||||
if dirName == "." || strings.Contains(dirName, "/") || strings.Contains(dirName, "\\") || strings.Contains(dirName, "..") {
|
||||
return false
|
||||
}
|
||||
for _, ch := range dirName {
|
||||
if !((ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || (ch >= '0' && ch <= '9') || ch == '_' || ch == '-' || ch == '.') {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// getRepoPhysicalPath 计算仓库任务的最终物理绝对路径
|
||||
func getRepoPhysicalPath(targetPath, dirName, sourceURL, branch string) string {
|
||||
if dirName == "." {
|
||||
return "" // 如果不追加目录,此逻辑不负责判断其根目录(共享的 scripts 目录)
|
||||
}
|
||||
finalDirName := dirName
|
||||
if finalDirName == "" {
|
||||
finalDirName = utils.GetRepoIdentifier(sourceURL, branch)
|
||||
}
|
||||
if finalDirName == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
basePath := targetPath
|
||||
if basePath == "" || basePath == constant.ScriptsDirPlaceholder {
|
||||
basePath = constant.ScriptsWorkDir
|
||||
} else if strings.HasPrefix(basePath, constant.ScriptsDirPlaceholder) {
|
||||
basePath = filepath.Join(constant.ScriptsWorkDir, strings.TrimPrefix(basePath, constant.ScriptsDirPlaceholder))
|
||||
} else if !filepath.IsAbs(basePath) {
|
||||
basePath = filepath.Join(constant.ScriptsWorkDir, basePath)
|
||||
}
|
||||
|
||||
fullPath := filepath.Join(basePath, finalDirName)
|
||||
absPath, err := filepath.Abs(fullPath)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return absPath
|
||||
}
|
||||
// CreateTask 创建任务
|
||||
// @Summary 创建任务
|
||||
// @Description 创建一个新的任务
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param body body vo.TaskCreateReq true "任务创建信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.TaskVO}
|
||||
// @Failure 400 {object} utils.Response
|
||||
// @Router /tasks [post]
|
||||
func (tc *TaskController) CreateTask(c *gin.Context) {
|
||||
var req vo.TaskCreateReq
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 普通任务需要命令
|
||||
if req.Type != constant.TaskTypeRepo && req.Command == "" {
|
||||
utils.BadRequest(c, "命令不能为空")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Schedule != "" {
|
||||
if err := tc.executorService.ValidateCron(req.Schedule); err != nil {
|
||||
utils.BadRequest(c, "无效的cron表达式: "+err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 转换为绝对路径(Agent 任务保持原样)
|
||||
workDir := req.WorkDir
|
||||
if req.AgentID == nil || *req.AgentID == "" {
|
||||
workDir = resolveWorkDir(req.WorkDir)
|
||||
}
|
||||
|
||||
var sourceID string
|
||||
// 如果是仓库同步任务,根据 URL 生成 SourceID 用于去重
|
||||
if req.Type == constant.TaskTypeRepo && req.Config != "" {
|
||||
var repoCfg struct {
|
||||
SourceURL string `json:"source_url"`
|
||||
Branch string `json:"branch"`
|
||||
RepoDirName string `json:"repo_dir_name"`
|
||||
TargetPath string `json:"target_path"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(req.Config), &repoCfg); err == nil && repoCfg.SourceURL != "" {
|
||||
if repoCfg.RepoDirName != "" {
|
||||
if !isValidDirName(repoCfg.RepoDirName) {
|
||||
utils.BadRequest(c, "自定义目录名只能包含字母、数字、下划线、短划线和点,不能只有点,且不能包含路径逻辑")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 如果配置了自定义名字,使用配置的名字。没有配置的话,使用以前的username_reponame
|
||||
if repoCfg.RepoDirName != "" {
|
||||
sourceID = "repo_" + repoCfg.RepoDirName
|
||||
} else {
|
||||
sourceID = "repo_" + utils.GetRepoIdentifier(repoCfg.SourceURL, repoCfg.Branch)
|
||||
}
|
||||
|
||||
// 校验 SourceID 是否已存在(任务唯一性)
|
||||
existingTask := tc.taskService.GetTaskBySourceID(sourceID)
|
||||
if existingTask != nil {
|
||||
utils.BadRequest(c, "当前任务已存在,请检查或更换仓库目录名称")
|
||||
return
|
||||
}
|
||||
|
||||
// 校验物理目录是否存在
|
||||
newAbsPath := getRepoPhysicalPath(repoCfg.TargetPath, repoCfg.RepoDirName, repoCfg.SourceURL, repoCfg.Branch)
|
||||
if newAbsPath != "" {
|
||||
if info, err := os.Stat(newAbsPath); err == nil && info.IsDir() {
|
||||
utils.BadRequest(c, "本地已存在同名仓库文件夹,请更换自定义目录名或清理残留文件")
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
param := tasks.TaskParam{
|
||||
Name: req.Name,
|
||||
Remark: req.Remark,
|
||||
Command: req.Command,
|
||||
PreCommand: req.PreCommand,
|
||||
PostCommand: req.PostCommand,
|
||||
Tags: req.Tags,
|
||||
Type: req.Type,
|
||||
Config: req.Config,
|
||||
Schedule: req.Schedule,
|
||||
Timeout: req.Timeout,
|
||||
WorkDir: workDir,
|
||||
CleanConfig: req.CleanConfig,
|
||||
Envs: req.Envs,
|
||||
Languages: req.Languages,
|
||||
AgentID: req.AgentID,
|
||||
TriggerType: req.TriggerType,
|
||||
RetryCount: req.RetryCount,
|
||||
RetryInterval: req.RetryInterval,
|
||||
RandomRange: req.RandomRange,
|
||||
SourceID: sourceID,
|
||||
PinType: req.PinType,
|
||||
Enabled: true,
|
||||
}
|
||||
|
||||
var task *models.Task
|
||||
// 去重逻辑:如果已存在相同 SourceID 的仓库任务,则改为更新
|
||||
if sourceID != "" {
|
||||
task = tc.taskService.GetTaskBySourceID(sourceID)
|
||||
if task != nil {
|
||||
task = tc.taskService.UpdateTask(task.ID, ¶m)
|
||||
}
|
||||
}
|
||||
|
||||
if task == nil {
|
||||
task = tc.taskService.CreateTask(¶m)
|
||||
}
|
||||
|
||||
// 如果是 Agent 任务,通知 Agent;否则添加到本地 cron
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*task.AgentID)
|
||||
} else {
|
||||
tc.executorService.AddCronTask(task)
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskVO(task))
|
||||
}
|
||||
|
||||
// BulkSaveTask 批量保存/导入任务配置(用于主节点下发同步)
|
||||
// @Summary 批量保存任务
|
||||
// @Description 批量导入任务配置,如果ID或同名存在则更新,不存在则创建
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Router /tasks/bulk_save [post]
|
||||
func (tc *TaskController) BulkSaveTask(c *gin.Context) {
|
||||
var reqs []vo.TaskVO
|
||||
|
||||
if err := c.ShouldBindJSON(&reqs); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
for _, req := range reqs {
|
||||
param := tasks.TaskParam{
|
||||
Name: req.Name,
|
||||
Remark: req.Remark,
|
||||
Command: req.Command,
|
||||
PreCommand: req.PreCommand,
|
||||
PostCommand: req.PostCommand,
|
||||
Tags: req.Tags,
|
||||
Type: req.Type,
|
||||
Config: req.Config,
|
||||
Schedule: req.Schedule,
|
||||
Timeout: req.Timeout,
|
||||
WorkDir: req.WorkDir,
|
||||
CleanConfig: req.CleanConfig,
|
||||
Envs: req.Envs,
|
||||
Languages: req.Languages,
|
||||
AgentID: req.AgentID,
|
||||
TriggerType: req.TriggerType,
|
||||
RetryCount: req.RetryCount,
|
||||
RetryInterval: req.RetryInterval,
|
||||
RandomRange: req.RandomRange,
|
||||
PinType: req.PinType,
|
||||
Enabled: req.Enabled,
|
||||
SourceID: "", // 不直接覆盖
|
||||
}
|
||||
|
||||
var existingTask *models.Task
|
||||
// 优先按 ID 匹配
|
||||
if req.ID != "" {
|
||||
existingTask = tc.taskService.GetTaskByID(req.ID)
|
||||
}
|
||||
// 如果 ID 没找到,尝试按 Name 匹配
|
||||
if existingTask == nil {
|
||||
var t models.Task
|
||||
res := database.DB.Where("name = ?", req.Name).First(&t)
|
||||
if res.Error == nil {
|
||||
existingTask = &t
|
||||
}
|
||||
}
|
||||
|
||||
var savedTask *models.Task
|
||||
if existingTask != nil {
|
||||
savedTask = tc.taskService.UpdateTask(existingTask.ID, ¶m)
|
||||
} else {
|
||||
savedTask = tc.taskService.CreateTask(¶m)
|
||||
// 如果原始有 ID,强制覆盖更新 ID 保持强同步一致性
|
||||
if req.ID != "" && savedTask != nil {
|
||||
database.DB.Model(savedTask).Update("id", req.ID)
|
||||
savedTask.ID = req.ID
|
||||
}
|
||||
}
|
||||
|
||||
// 如果是 Agent 任务,通知 Agent;否则添加到本地 cron
|
||||
if savedTask != nil {
|
||||
if savedTask.AgentID != nil && *savedTask.AgentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*savedTask.AgentID)
|
||||
} else {
|
||||
tc.executorService.AddCronTask(savedTask)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, nil)
|
||||
}
|
||||
|
||||
// GetTasks 获取任务列表
|
||||
// @Summary 获取任务列表
|
||||
// @Description 分页获取任务列表,支持按名称、Agent ID、标签、类型筛选
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param name query string false "任务名称"
|
||||
// @Param agent_id query string false "Agent ID"
|
||||
// @Param tags query string false "标签"
|
||||
// @Param type query string false "任务类型"
|
||||
// @Param page query int false "页码"
|
||||
// @Param page_size query int false "每页数量"
|
||||
// @Success 200 {object} utils.Response{data=utils.PaginationData{data=[]vo.TaskVO}}
|
||||
// @Router /tasks [get]
|
||||
func (tc *TaskController) GetTasks(c *gin.Context) {
|
||||
p := utils.ParsePagination(c)
|
||||
name := c.DefaultQuery("name", "")
|
||||
agentIDStr := c.DefaultQuery("agent_id", "")
|
||||
|
||||
tags := c.DefaultQuery("tags", "")
|
||||
taskType := c.DefaultQuery("type", "")
|
||||
|
||||
var agentID *string
|
||||
if agentIDStr != "" {
|
||||
agentID = &agentIDStr
|
||||
}
|
||||
|
||||
sortBy := c.DefaultQuery("sort_by", "")
|
||||
order := c.DefaultQuery("order", "")
|
||||
|
||||
tasks, total := tc.taskService.GetTasksWithPagination(p.Page, p.PageSize, name, agentID, tags, taskType, sortBy, order)
|
||||
utils.PaginatedResponse(c, vo.ToTaskVOListFromModels(tasks), total, p)
|
||||
}
|
||||
|
||||
// GetTask 获取任务详情
|
||||
// @Summary 获取任务详情
|
||||
// @Description 根据 ID 获取任务详情
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "任务ID"
|
||||
// @Success 200 {object} utils.Response{data=vo.TaskVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /tasks/{id} [get]
|
||||
func (tc *TaskController) GetTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
task := tc.taskService.GetTaskByID(id)
|
||||
if task == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskVO(task))
|
||||
}
|
||||
|
||||
// UpdateTask 更新任务
|
||||
// @Summary 更新任务
|
||||
// @Description 根据 ID 更新任务信息
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "任务ID"
|
||||
// @Param body body vo.TaskUpdateReq true "任务更新信息"
|
||||
// @Success 200 {object} utils.Response{data=vo.TaskVO}
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /tasks/{id} [put]
|
||||
func (tc *TaskController) UpdateTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取旧任务信息(用于判断 agent 变更)
|
||||
oldTask := tc.taskService.GetTaskByID(id)
|
||||
var oldAgentID *string
|
||||
if oldTask != nil {
|
||||
oldAgentID = oldTask.AgentID
|
||||
}
|
||||
|
||||
var req vo.TaskUpdateReq
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
if req.Schedule != "" {
|
||||
if err := tc.executorService.ValidateCron(req.Schedule); err != nil {
|
||||
utils.BadRequest(c, "无效的cron表达式: "+err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 转换为绝对路径(Agent 任务保持原样)
|
||||
workDir := req.WorkDir
|
||||
if req.AgentID == nil || *req.AgentID == "" {
|
||||
workDir = resolveWorkDir(req.WorkDir)
|
||||
}
|
||||
|
||||
var sourceID string
|
||||
if req.Type == constant.TaskTypeRepo && req.Config != "" {
|
||||
var repoCfg struct {
|
||||
SourceURL string `json:"source_url"`
|
||||
Branch string `json:"branch"`
|
||||
RepoDirName string `json:"repo_dir_name"`
|
||||
TargetPath string `json:"target_path"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(req.Config), &repoCfg); err == nil && repoCfg.SourceURL != "" {
|
||||
if repoCfg.RepoDirName != "" {
|
||||
if !isValidDirName(repoCfg.RepoDirName) {
|
||||
utils.BadRequest(c, "自定义目录名只能包含字母、数字、下划线、短划线和点,不能只有点,且不能包含路径逻辑")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 如果配置了自定义名字,使用配置的名字。没有配置的话,使用以前的username_reponame
|
||||
if repoCfg.RepoDirName != "" {
|
||||
sourceID = "repo_" + repoCfg.RepoDirName
|
||||
} else {
|
||||
sourceID = "repo_" + utils.GetRepoIdentifier(repoCfg.SourceURL, repoCfg.Branch)
|
||||
}
|
||||
|
||||
// 验证更新后的 SourceID 是否和别的任务冲突
|
||||
if sourceID != oldTask.SourceID {
|
||||
existingTask := tc.taskService.GetTaskBySourceID(sourceID)
|
||||
if existingTask != nil && existingTask.ID != oldTask.ID {
|
||||
utils.BadRequest(c, "当前任务已存在,请检查或更换仓库目录名称")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 计算新的物理路径
|
||||
newAbsPath := getRepoPhysicalPath(repoCfg.TargetPath, repoCfg.RepoDirName, repoCfg.SourceURL, repoCfg.Branch)
|
||||
|
||||
var oldAbsPath string
|
||||
if oldTask != nil && oldTask.Type == constant.TaskTypeRepo && oldTask.Config != "" {
|
||||
var oldCfg struct {
|
||||
SourceURL string `json:"source_url"`
|
||||
Branch string `json:"branch"`
|
||||
RepoDirName string `json:"repo_dir_name"`
|
||||
TargetPath string `json:"target_path"`
|
||||
}
|
||||
if json.Unmarshal([]byte(oldTask.Config), &oldCfg) == nil {
|
||||
oldAbsPath = getRepoPhysicalPath(oldCfg.TargetPath, oldCfg.RepoDirName, oldCfg.SourceURL, oldCfg.Branch)
|
||||
}
|
||||
}
|
||||
|
||||
// 如果路径发生了改变(或者是个全新计算的路径),并且新路径已存在,则报错拦截
|
||||
if newAbsPath != "" && newAbsPath != oldAbsPath {
|
||||
if info, err := os.Stat(newAbsPath); err == nil && info.IsDir() {
|
||||
utils.BadRequest(c, "目标目录在本地已存在同名文件夹,请更换目录名或清理残留文件")
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if oldTask != nil {
|
||||
sourceID = oldTask.SourceID
|
||||
}
|
||||
|
||||
param := tasks.TaskParam{
|
||||
Name: req.Name,
|
||||
Remark: req.Remark,
|
||||
Command: req.Command,
|
||||
PreCommand: req.PreCommand,
|
||||
PostCommand: req.PostCommand,
|
||||
Tags: req.Tags,
|
||||
Type: req.Type,
|
||||
Config: req.Config,
|
||||
Schedule: req.Schedule,
|
||||
Timeout: req.Timeout,
|
||||
WorkDir: workDir,
|
||||
CleanConfig: req.CleanConfig,
|
||||
Envs: req.Envs,
|
||||
Languages: req.Languages,
|
||||
AgentID: req.AgentID,
|
||||
TriggerType: req.TriggerType,
|
||||
RetryCount: req.RetryCount,
|
||||
RetryInterval: req.RetryInterval,
|
||||
RandomRange: req.RandomRange,
|
||||
SourceID: sourceID,
|
||||
PinType: req.PinType,
|
||||
Enabled: req.Enabled,
|
||||
}
|
||||
|
||||
task := tc.taskService.UpdateTask(id, ¶m)
|
||||
if task == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 处理任务调度
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
// Agent 任务:从本地 cron 移除,通知 Agent
|
||||
tc.executorService.RemoveCronTask(task.ID)
|
||||
tc.agentWSManager.BroadcastTasks(*task.AgentID)
|
||||
// 如果 agent 变更了,也通知旧 agent
|
||||
if oldAgentID != nil && *oldAgentID != "" && *oldAgentID != *task.AgentID {
|
||||
tc.agentWSManager.BroadcastTasks(*oldAgentID)
|
||||
}
|
||||
} else {
|
||||
// 本地任务
|
||||
if utils.DerefBool(task.Enabled, true) {
|
||||
tc.executorService.AddCronTask(task)
|
||||
} else {
|
||||
tc.executorService.RemoveCronTask(task.ID)
|
||||
}
|
||||
// 如果之前是 agent 任务,通知旧 agent 移除
|
||||
if oldAgentID != nil && *oldAgentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*oldAgentID)
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskVO(task))
|
||||
}
|
||||
|
||||
// DeleteTask 删除任务
|
||||
// @Summary 删除任务
|
||||
// @Description 根据 ID 删除任务
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param id path string true "任务ID"
|
||||
// @Success 200 {object} utils.Response
|
||||
// @Failure 404 {object} utils.Response
|
||||
// @Router /tasks/{id} [delete]
|
||||
func (tc *TaskController) DeleteTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取任务信息(用于通知 agent 和物理删除校验)
|
||||
task := tc.taskService.GetTaskByID(id)
|
||||
if task == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
agentID := task.AgentID
|
||||
deleteFiles := c.Query("delete_files") == "true"
|
||||
|
||||
// 如果需要删除物理文件且是仓库任务
|
||||
if deleteFiles && task.Type == constant.TaskTypeRepo {
|
||||
tc.deleteRepoPhysicalFiles(task)
|
||||
}
|
||||
|
||||
tc.executorService.RemoveCronTask(id)
|
||||
tc.executorService.GetScheduler().StopTask(id)
|
||||
|
||||
success := tc.taskService.DeleteTask(id)
|
||||
if !success {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 如果是 agent 任务,通知 agent
|
||||
if agentID != nil && *agentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*agentID)
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "删除成功")
|
||||
}
|
||||
|
||||
// deleteRepoPhysicalFiles 删除仓库关联的物理文件
|
||||
func (tc *TaskController) deleteRepoPhysicalFiles(task *models.Task) {
|
||||
if task.Type != constant.TaskTypeRepo {
|
||||
return
|
||||
}
|
||||
|
||||
logger.Infof("[Controller] 开始尝试物理删除任务关联文件: %s", task.Name)
|
||||
var repoCfg models.RepoConfig
|
||||
if err := json.Unmarshal([]byte(task.Config), &repoCfg); err != nil {
|
||||
logger.Errorf("[Controller] 解析任务配置失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
targetPath := repoCfg.TargetPath
|
||||
if targetPath == "" {
|
||||
// 如果 TargetPath 为空,调用系统的计算函数获取默认目录名
|
||||
repoId := utils.GetRepoIdentifier(repoCfg.SourceURL, repoCfg.Branch)
|
||||
if repoId != "" {
|
||||
targetPath = repoId
|
||||
logger.Infof("[Controller] TargetPath 为空,使用计算出的标识符: %s", targetPath)
|
||||
}
|
||||
}
|
||||
|
||||
if targetPath == "" || targetPath == constant.ScriptsDirPlaceholder {
|
||||
logger.Warnf("[Controller] 任务 %s 无法确定有效的物理删除路径,跳过", task.Name)
|
||||
return
|
||||
}
|
||||
|
||||
// 确定绝对路径
|
||||
scriptsDir, _ := filepath.Abs(constant.ScriptsWorkDir)
|
||||
fullPath := targetPath
|
||||
if strings.HasPrefix(targetPath, constant.ScriptsDirPlaceholder) {
|
||||
fullPath = filepath.Join(scriptsDir, strings.TrimPrefix(targetPath, constant.ScriptsDirPlaceholder))
|
||||
} else if !filepath.IsAbs(targetPath) {
|
||||
fullPath = filepath.Join(scriptsDir, targetPath)
|
||||
}
|
||||
|
||||
absTargetPath, _ := filepath.Abs(fullPath)
|
||||
logger.Infof("[Controller] 最终计算的绝对路径: %s, Scripts目录: %s", absTargetPath, scriptsDir)
|
||||
scriptsDir, _ = filepath.Abs(constant.ScriptsWorkDir)
|
||||
|
||||
// 安全检查:使用 Rel 判断路径关系
|
||||
rel, err := filepath.Rel(scriptsDir, absTargetPath)
|
||||
if err != nil {
|
||||
logger.Errorf("[Controller] 计算相对路径失败: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
// 必须是在 scripts 目录下(不以 .. 开头)且不能是 scripts 目录本身 (.)
|
||||
if rel != "." && !strings.HasPrefix(rel, "..") {
|
||||
err := os.RemoveAll(absTargetPath)
|
||||
if err != nil {
|
||||
logger.Errorf("[Controller] 物理删除文件夹失败: %s, 路径: %s, 错误: %v", task.Name, absTargetPath, err)
|
||||
} else {
|
||||
logger.Infof("[Controller] 已成功物理删除文件夹: %s, 路径: %s", task.Name, absTargetPath)
|
||||
}
|
||||
} else {
|
||||
logger.Warnf("[Controller] 拒绝物理删除安全目录之外的路径: %s", absTargetPath)
|
||||
}
|
||||
}
|
||||
|
||||
func (tc *TaskController) BatchDeleteTasks(c *gin.Context) {
|
||||
var req struct {
|
||||
IDs []string `json:"ids" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// 收集涉及到的 AgentID
|
||||
agentIDs := make(map[string]struct{})
|
||||
for _, id := range req.IDs {
|
||||
// 获取任务信息
|
||||
task := tc.taskService.GetTaskByID(id)
|
||||
if task != nil {
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
agentIDs[*task.AgentID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
// 移除 cron 调度
|
||||
tc.executorService.RemoveCronTask(id)
|
||||
tc.executorService.GetScheduler().StopTask(id)
|
||||
}
|
||||
|
||||
// 执行批量删除
|
||||
count := tc.taskService.BatchDeleteTasks(req.IDs)
|
||||
|
||||
// 通知受影响的 Agent
|
||||
for agentID := range agentIDs {
|
||||
tc.agentWSManager.BroadcastTasks(agentID)
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{"count": count})
|
||||
}
|
||||
|
||||
// BatchDeleteByQuery 根据查询条件批量删除任务
|
||||
// @Summary 根据查询条件批量删除任务
|
||||
// @Description 根据查询条件批量删除匹配的所有任务
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param name query string false "任务名称关键词"
|
||||
// @Param tags query string false "标签关键词"
|
||||
// @Param type query string false "任务类型"
|
||||
// @Param agent_id query string false "执行位置(节点ID)"
|
||||
// @Success 200 {object} utils.Response{data=map[string]int}
|
||||
// @Failure 401 {object} utils.Response "未授权"
|
||||
// @Router /tasks/batch-by-query [delete]
|
||||
func (tc *TaskController) BatchDeleteByQuery(c *gin.Context) {
|
||||
name := c.Query("name")
|
||||
agentIDStr := c.Query("agent_id")
|
||||
tags := c.Query("tags")
|
||||
taskType := c.Query("type")
|
||||
|
||||
var agentID *string
|
||||
if agentIDStr != "" {
|
||||
agentID = &agentIDStr
|
||||
}
|
||||
|
||||
tasks, _ := tc.taskService.GetTasksWithPagination(1, 999999, name, agentID, tags, taskType, "", "")
|
||||
if len(tasks) == 0 {
|
||||
utils.Success(c, gin.H{"count": 0})
|
||||
return
|
||||
}
|
||||
|
||||
var ids []string
|
||||
agentIDs := make(map[string]struct{})
|
||||
for _, task := range tasks {
|
||||
ids = append(ids, task.ID)
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
agentIDs[*task.AgentID] = struct{}{}
|
||||
}
|
||||
// 移除 cron 调度
|
||||
tc.executorService.RemoveCronTask(task.ID)
|
||||
tc.executorService.GetScheduler().StopTask(task.ID)
|
||||
}
|
||||
|
||||
// 执行批量删除
|
||||
count := tc.taskService.BatchDeleteTasks(ids)
|
||||
|
||||
// 通知受影响的 Agent
|
||||
for aID := range agentIDs {
|
||||
tc.agentWSManager.BroadcastTasks(aID)
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{"count": count})
|
||||
}
|
||||
|
||||
// StopTask 停止任务
|
||||
// @Summary 停止任务
|
||||
// @Description 根据运行日志 ID 停止正在执行的任务
|
||||
// @Tags 任务管理
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param logID path string true "运行日志ID"
|
||||
// @Success 200 {object} utils.Response
|
||||
// @Failure 400 {object} utils.Response
|
||||
// @Router /tasks/stop/{logID} [post]
|
||||
func (tc *TaskController) StopTask(c *gin.Context) {
|
||||
logID := c.Param("logID")
|
||||
if logID == "" {
|
||||
utils.BadRequest(c, "无效的日志ID")
|
||||
return
|
||||
}
|
||||
|
||||
err := tc.executorService.StopTaskExecution(logID)
|
||||
if err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.SuccessMsg(c, "停止请求已发送")
|
||||
}
|
||||
|
||||
// GetTags 获取所有任务标签
|
||||
// @Summary 获取所有任务标签
|
||||
// @Description 获取系统中所有任务已使用的唯一标签列表
|
||||
// @Tags 任务管理
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Success 200 {object} utils.Response{data=[]string}
|
||||
// @Router /tasks/tags [get]
|
||||
func (tc *TaskController) GetTags(c *gin.Context) {
|
||||
tags, err := tc.taskService.GetAllTags()
|
||||
if err != nil {
|
||||
utils.ServerError(c, err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(c, tags)
|
||||
}
|
||||
|
||||
// SyncRepoTasks 增量同步仓库任务状态(供本地 reposync 进程调用)
|
||||
func (tc *TaskController) SyncRepoTasks(c *gin.Context) {
|
||||
var req struct {
|
||||
RepoID string `json:"repo_id"`
|
||||
UpsertedIDs []string `json:"upserted_ids"`
|
||||
DeletedIDs []string `json:"deleted_ids"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
tc.executorService.SyncRepoTasks(req.UpsertedIDs, req.DeletedIDs)
|
||||
utils.SuccessMsg(c, "增量同步成功")
|
||||
}
|
||||
|
||||
// ToggleTask 切换任务启用/禁用状态
|
||||
func (tc *TaskController) ToggleTask(c *gin.Context) {
|
||||
id := c.Param("id")
|
||||
if id == "" {
|
||||
utils.BadRequest(c, "无效的任务ID")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
task := tc.taskService.GetTaskByID(id)
|
||||
if task == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 获取旧 AgentID
|
||||
var oldAgentID *string
|
||||
oldAgentID = task.AgentID
|
||||
|
||||
// 构造更新参数,仅修改 Enabled
|
||||
param := tasks.TaskParam{
|
||||
Name: task.Name,
|
||||
Remark: task.Remark,
|
||||
Command: string(task.Command),
|
||||
PreCommand: string(task.PreCommand),
|
||||
PostCommand: string(task.PostCommand),
|
||||
Tags: task.Tags,
|
||||
Type: task.Type,
|
||||
Config: string(task.Config),
|
||||
Schedule: task.Schedule,
|
||||
Timeout: task.Timeout,
|
||||
WorkDir: task.WorkDir,
|
||||
CleanConfig: task.CleanConfig,
|
||||
Envs: string(task.Envs),
|
||||
Languages: task.Languages,
|
||||
AgentID: task.AgentID,
|
||||
TriggerType: task.TriggerType,
|
||||
RetryCount: task.RetryCount,
|
||||
RetryInterval: task.RetryInterval,
|
||||
RandomRange: task.RandomRange,
|
||||
SourceID: task.SourceID,
|
||||
PinType: task.PinType,
|
||||
Enabled: req.Enabled,
|
||||
}
|
||||
|
||||
updatedTask := tc.taskService.UpdateTask(id, ¶m)
|
||||
if updatedTask == nil {
|
||||
utils.NotFound(c, "任务不存在")
|
||||
return
|
||||
}
|
||||
|
||||
// 处理调度器更新
|
||||
if updatedTask.AgentID != nil && *updatedTask.AgentID != "" {
|
||||
tc.executorService.RemoveCronTask(updatedTask.ID)
|
||||
tc.agentWSManager.BroadcastTasks(*updatedTask.AgentID)
|
||||
} else {
|
||||
if req.Enabled {
|
||||
tc.executorService.AddCronTask(updatedTask)
|
||||
} else {
|
||||
tc.executorService.RemoveCronTask(updatedTask.ID)
|
||||
}
|
||||
if oldAgentID != nil && *oldAgentID != "" {
|
||||
tc.agentWSManager.BroadcastTasks(*oldAgentID)
|
||||
}
|
||||
}
|
||||
|
||||
utils.Success(c, vo.ToTaskVO(updatedTask))
|
||||
}
|
||||
@@ -0,0 +1,476 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
|
||||
"github.com/creack/pty"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
"golang.org/x/text/encoding/simplifiedchinese"
|
||||
"golang.org/x/text/transform"
|
||||
)
|
||||
|
||||
type TerminalController struct {
|
||||
envService *services.EnvService
|
||||
}
|
||||
|
||||
func NewTerminalController(envService *services.EnvService) *TerminalController {
|
||||
return &TerminalController{
|
||||
envService: envService,
|
||||
}
|
||||
}
|
||||
|
||||
var upgrader = websocket.Upgrader{
|
||||
CheckOrigin: utils.CheckWSOrigin,
|
||||
}
|
||||
|
||||
// toUTF8 将可能是 GBK 编码的字节转换为 UTF-8
|
||||
func toUTF8(data []byte) string {
|
||||
if utf8.Valid(data) {
|
||||
return string(data)
|
||||
}
|
||||
// 尝试从 GBK 转换
|
||||
reader := transform.NewReader(
|
||||
bufio.NewReader(
|
||||
&byteReader{data: data},
|
||||
),
|
||||
simplifiedchinese.GBK.NewDecoder(),
|
||||
)
|
||||
result, err := io.ReadAll(reader)
|
||||
if err != nil {
|
||||
return string(data)
|
||||
}
|
||||
return string(result)
|
||||
}
|
||||
|
||||
type byteReader struct {
|
||||
data []byte
|
||||
pos int
|
||||
}
|
||||
|
||||
func (r *byteReader) Read(p []byte) (n int, err error) {
|
||||
if r.pos >= len(r.data) {
|
||||
return 0, io.EOF
|
||||
}
|
||||
n = copy(p, r.data[r.pos:])
|
||||
r.pos += n
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (tc *TerminalController) HandleWebSocket(c *gin.Context) {
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// 演示模式下禁用终端
|
||||
if constant.DemoMode {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("\r\n\033[1;33m[演示模式] 终端功能已禁用\033[0m\r\n"))
|
||||
return
|
||||
}
|
||||
|
||||
// Windows 使用 pipe 模式,Unix 使用 PTY 模式
|
||||
userID := c.GetString("userID")
|
||||
if userID == "" {
|
||||
userID = "1" // 兜底
|
||||
}
|
||||
if runtime.GOOS == "windows" {
|
||||
tc.handlePipeMode(conn, userID)
|
||||
} else {
|
||||
tc.handlePtyMode(conn, userID)
|
||||
}
|
||||
}
|
||||
|
||||
// handlePtyMode 使用 PTY 处理终端(Unix/macOS)
|
||||
func (tc *TerminalController) handlePtyMode(conn *websocket.Conn, userID string) {
|
||||
conn.SetReadLimit(constant.MaxMessageSize)
|
||||
conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
return nil
|
||||
})
|
||||
|
||||
// 发送 PTY 模式标识
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("__PTY_MODE__"))
|
||||
|
||||
cmd := utils.NewShellCmd()
|
||||
|
||||
if absDir, err := filepath.Abs(constant.ScriptsWorkDir); err == nil {
|
||||
cmd.Dir = absDir
|
||||
}
|
||||
|
||||
cmd.Env = tc.buildTerminalEnv(userID, "TERM=xterm-256color")
|
||||
|
||||
ptmx, err := pty.Start(cmd)
|
||||
if err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error starting shell: "+err.Error()))
|
||||
return
|
||||
}
|
||||
defer ptmx.Close()
|
||||
|
||||
pty.Setsize(ptmx, &pty.Winsize{Rows: 24, Cols: 80})
|
||||
|
||||
var wg sync.WaitGroup
|
||||
var connMu sync.Mutex
|
||||
|
||||
writeMessage := func(data []byte) {
|
||||
connMu.Lock()
|
||||
defer connMu.Unlock()
|
||||
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
conn.WriteMessage(websocket.TextMessage, data)
|
||||
}
|
||||
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
buf := make([]byte, 4096)
|
||||
var remainder []byte
|
||||
for {
|
||||
n, err := ptmx.Read(buf)
|
||||
if n > 0 {
|
||||
chunk := append(remainder, buf[:n]...)
|
||||
lastSafe := len(chunk)
|
||||
for i := len(chunk); i > 0 && i > len(chunk)-4; i-- {
|
||||
if utf8.RuneStart(chunk[i-1]) {
|
||||
if !utf8.FullRune(chunk[i-1 : len(chunk)]) {
|
||||
lastSafe = i - 1
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
safe := chunk[:lastSafe]
|
||||
|
||||
if len(safe) == 0 && len(chunk) >= 4 {
|
||||
safe = chunk
|
||||
remainder = nil
|
||||
} else {
|
||||
remainder = make([]byte, len(chunk[lastSafe:]))
|
||||
copy(remainder, chunk[lastSafe:])
|
||||
}
|
||||
|
||||
if len(safe) > 0 {
|
||||
text := toUTF8(safe)
|
||||
writeMessage([]byte(text))
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
if len(remainder) > 0 {
|
||||
writeMessage([]byte(toUTF8(remainder)))
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
// 启动 ping 协程
|
||||
pingDone := make(chan struct{})
|
||||
go func() {
|
||||
ticker := time.NewTicker(constant.PingPeriod)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
connMu.Lock()
|
||||
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
||||
connMu.Unlock()
|
||||
return
|
||||
}
|
||||
connMu.Unlock()
|
||||
case <-pingDone:
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
_, message, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
|
||||
// 处理调整窗口大小的消息
|
||||
if len(message) > 0 && message[0] == '{' {
|
||||
var resizeMsg struct {
|
||||
Type string `json:"type"`
|
||||
Rows uint16 `json:"rows"`
|
||||
Cols uint16 `json:"cols"`
|
||||
}
|
||||
if err := json.Unmarshal(message, &resizeMsg); err == nil && resizeMsg.Type == "resize" {
|
||||
pty.Setsize(ptmx, &pty.Winsize{Rows: resizeMsg.Rows, Cols: resizeMsg.Cols})
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := ptmx.Write(message); err != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
close(pingDone)
|
||||
cmd.Process.Kill()
|
||||
cmd.Wait()
|
||||
ptmx.Close() // Force close PTY to interrupt the blocking ptmx.Read() in the goroutine
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// handlePipeMode 使用 pipe 处理终端(Windows)
|
||||
func (tc *TerminalController) handlePipeMode(conn *websocket.Conn, userID string) {
|
||||
conn.SetReadLimit(constant.MaxMessageSize)
|
||||
conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
conn.SetPongHandler(func(string) error {
|
||||
conn.SetReadDeadline(time.Now().Add(constant.PongWait))
|
||||
return nil
|
||||
})
|
||||
|
||||
// 发送 pipe 模式标识
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("__PIPE_MODE__"))
|
||||
|
||||
cmd := utils.NewShellCmd()
|
||||
|
||||
if absDir, err := filepath.Abs(constant.ScriptsWorkDir); err == nil {
|
||||
cmd.Dir = absDir
|
||||
}
|
||||
|
||||
// 注入环境变量
|
||||
cmd.Env = tc.buildTerminalEnv(userID)
|
||||
|
||||
stdin, err := cmd.StdinPipe()
|
||||
if err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
stdout, err := cmd.StdoutPipe()
|
||||
if err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
stderr, err := cmd.StderrPipe()
|
||||
if err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
conn.WriteMessage(websocket.TextMessage, []byte("Error: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
var connMu sync.Mutex
|
||||
|
||||
writeMessage := func(data []byte) {
|
||||
connMu.Lock()
|
||||
defer connMu.Unlock()
|
||||
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
conn.WriteMessage(websocket.TextMessage, data)
|
||||
}
|
||||
|
||||
readOutput := func(reader io.Reader) {
|
||||
defer wg.Done()
|
||||
defer func() { recover() }()
|
||||
buf := make([]byte, 4096)
|
||||
var remainder []byte
|
||||
for {
|
||||
n, err := reader.Read(buf)
|
||||
if n > 0 {
|
||||
chunk := append(remainder, buf[:n]...)
|
||||
lastSafe := len(chunk)
|
||||
for i := len(chunk); i > 0 && i > len(chunk)-4; i-- {
|
||||
if utf8.RuneStart(chunk[i-1]) {
|
||||
if !utf8.FullRune(chunk[i-1 : len(chunk)]) {
|
||||
lastSafe = i - 1
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
safe := chunk[:lastSafe]
|
||||
|
||||
if len(safe) == 0 && len(chunk) >= 4 {
|
||||
safe = chunk
|
||||
remainder = nil
|
||||
} else {
|
||||
remainder = make([]byte, len(chunk[lastSafe:]))
|
||||
copy(remainder, chunk[lastSafe:])
|
||||
}
|
||||
|
||||
if len(safe) > 0 {
|
||||
text := toUTF8(safe)
|
||||
writeMessage([]byte(text))
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
if len(remainder) > 0 {
|
||||
writeMessage([]byte(toUTF8(remainder)))
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
wg.Add(2)
|
||||
go readOutput(stdout)
|
||||
go readOutput(stderr)
|
||||
|
||||
// 启动 ping 协程
|
||||
pingDone := make(chan struct{})
|
||||
go func() {
|
||||
ticker := time.NewTicker(constant.PingPeriod)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
connMu.Lock()
|
||||
conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
|
||||
if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
||||
connMu.Unlock()
|
||||
return
|
||||
}
|
||||
connMu.Unlock()
|
||||
case <-pingDone:
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
_, message, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
// 过滤调整窗口大小的消息(Windows Pipe 模式不支持调整尺寸,需过滤掉避免写入 stdin)
|
||||
if len(message) > 0 && message[0] == '{' {
|
||||
var resizeMsg struct {
|
||||
Type string `json:"type"`
|
||||
}
|
||||
if err := json.Unmarshal(message, &resizeMsg); err == nil && resizeMsg.Type == "resize" {
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := stdin.Write(message); err != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
close(pingDone)
|
||||
cmd.Process.Kill()
|
||||
cmd.Wait()
|
||||
|
||||
if stdinCloser, ok := stdin.(io.Closer); ok {
|
||||
stdinCloser.Close()
|
||||
}
|
||||
if stdoutCloser, ok := stdout.(io.Closer); ok {
|
||||
stdoutCloser.Close()
|
||||
}
|
||||
if stderrCloser, ok := stderr.(io.Closer); ok {
|
||||
stderrCloser.Close()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// ExecuteShellCommand 执行单个命令并返回结果
|
||||
func (tc *TerminalController) ExecuteShellCommand(c *gin.Context) {
|
||||
// 演示模式下禁止执行命令
|
||||
if constant.DemoMode {
|
||||
utils.BadRequest(c, "演示模式下不能执行命令")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Command string `json:"command" binding:"required"`
|
||||
}
|
||||
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
cmd := utils.NewShellCommandCmd(req.Command)
|
||||
userID := c.GetString("userID")
|
||||
if userID == "" {
|
||||
userID = "1" // 与 WebSocket 终端保持一致,保留原有兜底行为
|
||||
}
|
||||
cmd.Env = tc.buildTerminalEnv(userID)
|
||||
output, err := cmd.CombinedOutput()
|
||||
|
||||
if err != nil {
|
||||
utils.Success(c, gin.H{
|
||||
"output": string(output),
|
||||
"error": err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(c, gin.H{
|
||||
"output": string(output),
|
||||
})
|
||||
}
|
||||
|
||||
func (tc *TerminalController) buildTerminalEnv(userID string, extraEnvs ...string) []string {
|
||||
env := os.Environ()
|
||||
env = append(env, extraEnvs...)
|
||||
|
||||
// 注入 taskpool 命令运行时路径与配置,保证在终端里手动执行 taskpool 子命令时
|
||||
// 仍然能连接到与主服务一致的数据库,而不是回退到默认 sqlite。
|
||||
if absBinDir, err := filepath.Abs(filepath.Join(constant.DataDir, "bin")); err == nil {
|
||||
pathStr := absBinDir + string(os.PathListSeparator) + os.Getenv("PATH")
|
||||
env = append(env, "PATH="+pathStr)
|
||||
}
|
||||
env = append(env, utils.BuildRuntimeProcessEnv()...)
|
||||
|
||||
// 注入环境变量(支持同名合并)
|
||||
if tc.envService != nil {
|
||||
envVars := tc.envService.GetFormattedEnvVarsByUserID(userID)
|
||||
env = append(env, envVars...)
|
||||
}
|
||||
|
||||
// 为 Docker 环境或二进制版本注入所有 mise 已安装 Node 的全局依赖路径到 NODE_PATH (Issue-90)
|
||||
if !utils.IsInDocker() || (!strings.Contains(os.Args[0], "go-build") && !strings.Contains(os.Args[0], "tmp")) {
|
||||
versions, _ := utils.ListMiseInstalledVersions("node")
|
||||
var nodePaths []string
|
||||
for _, v := range versions {
|
||||
if p := utils.GetMiseNodePath(v); p != "" {
|
||||
nodePaths = append(nodePaths, p)
|
||||
}
|
||||
}
|
||||
if len(nodePaths) > 0 {
|
||||
sep := ":"
|
||||
if runtime.GOOS == "windows" {
|
||||
sep = ";"
|
||||
}
|
||||
env = append(env, "NODE_PATH="+strings.Join(nodePaths, sep))
|
||||
}
|
||||
}
|
||||
|
||||
return env
|
||||
}
|
||||
|
||||
// GetCommands 获取所有可用的 cmd 列表及说明
|
||||
func (tc *TerminalController) GetCommands(c *gin.Context) {
|
||||
var cmds []map[string]string
|
||||
for _, cmdInfo := range constant.Commands {
|
||||
cmds = append(cmds, map[string]string{
|
||||
"name": cmdInfo.Name,
|
||||
"description": cmdInfo.Description,
|
||||
})
|
||||
}
|
||||
utils.Success(c, cmds)
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
package controllers
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type WebUIController struct {
|
||||
webuiService *services.WebUIService
|
||||
}
|
||||
|
||||
func NewWebUIController(webuiService *services.WebUIService) *WebUIController {
|
||||
return &WebUIController{
|
||||
webuiService: webuiService,
|
||||
}
|
||||
}
|
||||
|
||||
// GetWebUIs 获取所有WebUI
|
||||
func (c *WebUIController) GetWebUIs(ctx *gin.Context) {
|
||||
webuis, err := c.webuiService.GetWebUIs()
|
||||
if err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
utils.Success(ctx, webuis)
|
||||
}
|
||||
|
||||
// UploadWebUI 上传新WebUI
|
||||
func (c *WebUIController) UploadWebUI(ctx *gin.Context) {
|
||||
file, err := ctx.FormFile("file")
|
||||
if err != nil {
|
||||
utils.BadRequest(ctx, "获取上传文件失败")
|
||||
return
|
||||
}
|
||||
|
||||
// 临时保存上传的文件到挂载目录,避免 /tmp 跨分区移动或权限问题
|
||||
tmpDir := filepath.Join(constant.DataDir, "tmp")
|
||||
os.MkdirAll(tmpDir, 0755)
|
||||
tmpFile := filepath.Join(tmpDir, file.Filename)
|
||||
|
||||
if err := ctx.SaveUploadedFile(file, tmpFile); err != nil {
|
||||
utils.ServerError(ctx, "保存临时文件失败")
|
||||
return
|
||||
}
|
||||
defer os.Remove(tmpFile) // 自动清理临时文件
|
||||
|
||||
webuiName, err := c.webuiService.ExtractWebUI(tmpFile)
|
||||
if err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"message": "WebUI上传成功", "webui": webuiName})
|
||||
}
|
||||
|
||||
// SetActiveWebUI 切换活动WebUI
|
||||
func (c *WebUIController) SetActiveWebUI(ctx *gin.Context) {
|
||||
var req struct {
|
||||
Name string `json:"name" binding:"required"`
|
||||
}
|
||||
if err := ctx.ShouldBindJSON(&req); err != nil {
|
||||
utils.BadRequest(ctx, "无效的请求参数")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.webuiService.SetActiveWebUI(req.Name); err != nil {
|
||||
utils.ServerError(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"message": "WebUI已切换成功,部分页面可能需要刷新"})
|
||||
}
|
||||
|
||||
// DeleteWebUI 删除自定义WebUI
|
||||
func (c *WebUIController) DeleteWebUI(ctx *gin.Context) {
|
||||
name := ctx.Param("name")
|
||||
if name == "" {
|
||||
utils.BadRequest(ctx, "未提供WebUI名称")
|
||||
return
|
||||
}
|
||||
|
||||
if err := c.webuiService.DeleteWebUI(name); err != nil {
|
||||
utils.BadRequest(ctx, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
utils.Success(ctx, gin.H{"message": "WebUI已删除"})
|
||||
}
|
||||
@@ -0,0 +1,130 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/systime"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/driver/mysql"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
gormlogger "gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
var DB *gorm.DB
|
||||
var DBConfig *Config
|
||||
|
||||
type Config struct {
|
||||
Type string // sqlite, mysql, postgres
|
||||
Host string
|
||||
Port int
|
||||
User string
|
||||
Password string
|
||||
DBName string
|
||||
Path string // for sqlite
|
||||
DSN string // for mysql/mariadb unix socket or custom dsn
|
||||
SSLMode string // postgres: disable/require/verify-ca/verify-full; mysql: true/skip-verify
|
||||
}
|
||||
|
||||
func Init(cfg *Config) error {
|
||||
var err error
|
||||
DBConfig = cfg
|
||||
// 设置东八区时区
|
||||
loc := systime.CST
|
||||
time.Local = loc
|
||||
|
||||
dsn, err := buildDSN(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var dialector gorm.Dialector
|
||||
|
||||
switch cfg.Type {
|
||||
case "sqlite":
|
||||
dialector = sqlite.Open(dsn)
|
||||
case "mysql":
|
||||
dialector = mysql.Open(dsn)
|
||||
case "postgres":
|
||||
dialector = postgres.Open(dsn)
|
||||
default:
|
||||
return fmt.Errorf("unsupported database type: %s", cfg.Type)
|
||||
}
|
||||
|
||||
newLogger := gormlogger.New(
|
||||
log.New(os.Stdout, "\r\n", log.LstdFlags), // io writer
|
||||
gormlogger.Config{
|
||||
SlowThreshold: time.Millisecond * 500, // 慢 SQL 阈值,默认是 200ms,这里改为 500ms
|
||||
LogLevel: gormlogger.Warn, // 日志级别
|
||||
IgnoreRecordNotFoundError: true, // 忽略 ErrRecordNotFound(找不到记录)错误
|
||||
Colorful: true, // 禁用彩色打印
|
||||
},
|
||||
)
|
||||
|
||||
DB, err = gorm.Open(dialector, &gorm.Config{
|
||||
Logger: newLogger,
|
||||
NowFunc: func() time.Time {
|
||||
return time.Now().In(loc)
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to connect database: %w", err)
|
||||
}
|
||||
|
||||
logger.Infof("[Database] 已连接 %s 数据库 (时区: Asia/Shanghai)", cfg.Type)
|
||||
|
||||
// SQLite 特殊优化:开启 WAL 模式,提升并发性能
|
||||
if cfg.Type == "sqlite" {
|
||||
sqlDB, _ := DB.DB()
|
||||
if sqlDB != nil {
|
||||
sqlDB.SetMaxOpenConns(1) // SQLite 只允许单写连接
|
||||
sqlDB.Exec("PRAGMA journal_mode=WAL")
|
||||
sqlDB.Exec("PRAGMA synchronous=NORMAL")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func AutoMigrate(models ...interface{}) error {
|
||||
return DB.AutoMigrate(models...)
|
||||
}
|
||||
|
||||
func GetDB() *gorm.DB {
|
||||
return DB
|
||||
}
|
||||
|
||||
func buildDSN(cfg *Config) (string, error) {
|
||||
switch cfg.Type {
|
||||
case "sqlite":
|
||||
return cfg.Path + "?_busy_timeout=5000", nil
|
||||
case "mysql":
|
||||
if cfg.DSN != "" {
|
||||
return cfg.DSN, nil
|
||||
}
|
||||
dsn := fmt.Sprintf("%s:%s@tcp(%s:%d)/%s?charset=utf8mb4&parseTime=True&loc=Asia%%2FShanghai",
|
||||
cfg.User, cfg.Password, cfg.Host, cfg.Port, cfg.DBName)
|
||||
if cfg.SSLMode != "" {
|
||||
dsn += "&tls=" + cfg.SSLMode
|
||||
}
|
||||
return dsn, nil
|
||||
case "postgres":
|
||||
if cfg.DSN != "" {
|
||||
return cfg.DSN, nil
|
||||
}
|
||||
sslMode := cfg.SSLMode
|
||||
if sslMode == "" {
|
||||
sslMode = "disable"
|
||||
}
|
||||
dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s sslmode=%s TimeZone=Asia/Shanghai",
|
||||
cfg.Host, cfg.Port, cfg.User, cfg.Password, cfg.DBName, sslMode)
|
||||
return dsn, nil
|
||||
default:
|
||||
return "", fmt.Errorf("unsupported database type: %s", cfg.Type)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,304 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"crypto/md5"
|
||||
"encoding/hex"
|
||||
"reflect"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/rs/xid"
|
||||
)
|
||||
|
||||
var allModels = []interface{}{
|
||||
&models.AppLog{},
|
||||
&models.User{},
|
||||
&models.Task{},
|
||||
&models.TaskLog{},
|
||||
&models.Script{},
|
||||
&models.EnvironmentVariable{},
|
||||
&models.Setting{},
|
||||
&models.SendStats{},
|
||||
&models.Dependency{},
|
||||
&models.Agent{},
|
||||
&models.AgentToken{},
|
||||
&models.Language{},
|
||||
&models.NotifyWay{},
|
||||
&models.NotifyBinding{},
|
||||
&models.DataRelation{},
|
||||
&models.DataStorage{},
|
||||
&models.InterconnectNode{},
|
||||
}
|
||||
|
||||
func Migrate() error {
|
||||
// 1. 自动指纹识别,大幅提升远程数据库启动进度
|
||||
sig := getModelSignature(allModels)
|
||||
if DB.Migrator().HasTable(&models.Setting{}) {
|
||||
var sigSetting models.Setting
|
||||
res := DB.Where(&models.Setting{Section: constant.SectionSystem, Key: constant.KeySchemaSignature}).Limit(1).Find(&sigSetting)
|
||||
if res.RowsAffected > 0 && string(sigSetting.Value) == sig {
|
||||
logger.Info("[Database] 模型指纹一致,跳过自动表结构同步")
|
||||
|
||||
// 即使表结构一致,也要执行后置数据迁移(内部有幂等检查),防止有漏网之鱼
|
||||
logger.Info("[Database] 正在执行后置数据迁移...")
|
||||
if err := postMigrations(); err != nil {
|
||||
logger.Warnf("[Database] 后置数据迁移警告: %v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// 执行前置结构迁移
|
||||
logger.Info("[Database] 正在执行前置结构迁移与表结构同步...")
|
||||
if err := preMigrations(); err != nil {
|
||||
logger.Warnf("[Database] 前置结构迁移警告: %v", err)
|
||||
}
|
||||
|
||||
logger.Infof("[Database] 正在同步 %d 个数据模型的表结构...", len(allModels))
|
||||
if err := AutoMigrate(allModels...); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 执行后置数据迁移,依赖完整的表结构
|
||||
logger.Info("[Database] 正在执行后置数据迁移...")
|
||||
if err := postMigrations(); err != nil {
|
||||
logger.Warnf("[Database] 后置数据迁移警告: %v", err)
|
||||
}
|
||||
|
||||
// 3. 更新指纹记录
|
||||
if DB.Migrator().HasTable(&models.Setting{}) {
|
||||
var sigSetting models.Setting
|
||||
res := DB.Where(&models.Setting{Section: constant.SectionSystem, Key: constant.KeySchemaSignature}).Limit(1).Find(&sigSetting)
|
||||
if res.RowsAffected > 0 {
|
||||
DB.Model(&sigSetting).Update("value", models.BigText(sig))
|
||||
} else {
|
||||
DB.Create(&models.Setting{
|
||||
ID: constant.IDSchemaSignature,
|
||||
Section: constant.SectionSystem,
|
||||
Key: constant.KeySchemaSignature,
|
||||
Value: models.BigText(sig),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// getModelSignature 生成数据模型的结构指纹
|
||||
func getModelSignature(models []interface{}) string {
|
||||
var sb strings.Builder
|
||||
// 包含表前缀,确保前缀变更时也能触发迁移
|
||||
sb.WriteString(constant.TablePrefix)
|
||||
for _, m := range models {
|
||||
t := reflect.TypeOf(m)
|
||||
if t.Kind() == reflect.Ptr {
|
||||
t = t.Elem()
|
||||
}
|
||||
sb.WriteString(t.Name())
|
||||
for i := 0; i < t.NumField(); i++ {
|
||||
f := t.Field(i)
|
||||
if f.Anonymous {
|
||||
continue
|
||||
}
|
||||
sb.WriteString(f.Name)
|
||||
sb.WriteString(f.Type.String())
|
||||
sb.WriteString(f.Tag.Get("gorm"))
|
||||
}
|
||||
}
|
||||
hash := md5.Sum([]byte(sb.String()))
|
||||
return hex.EncodeToString(hash[:])
|
||||
}
|
||||
|
||||
// preMigrations 前置结构迁移,处理 AutoMigrate 无法自动解决的变更
|
||||
func preMigrations() error {
|
||||
// 检查 ql_tokens 表是否存在
|
||||
if DB.Migrator().HasTable(constant.TableMigrateQlTokens) {
|
||||
// 如果 code 列存在,且 token 列不存在,则重命名
|
||||
if DB.Migrator().HasColumn(&models.AgentToken{}, constant.ColumnMigrateQlTokenCode) {
|
||||
if err := DB.Migrator().RenameColumn(&models.AgentToken{}, constant.ColumnMigrateQlTokenCode, constant.ColumnMigrateQlTokenToken); err != nil {
|
||||
logger.Debugf("[Database] 重命名 ql_tokens.code 失败: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
// 移除 deps 表中的 type 字段(如果存在)
|
||||
if DB.Migrator().HasColumn(&models.Dependency{}, constant.ColumnMigrateDependencyType) {
|
||||
if err := DB.Migrator().DropColumn(&models.Dependency{}, constant.ColumnMigrateDependencyType); err != nil {
|
||||
logger.Debugf("[Database] 移除 deps.type 字段失败: %v", err)
|
||||
} else {
|
||||
logger.Infof("[Database] 已成功移除 deps 表中的 type 字段")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// postMigrations 数据后置迁移,用于需要依赖完整表结构的数据搬运
|
||||
func postMigrations() error {
|
||||
// 迁移任务标签到通用的数据关联表中
|
||||
migrateTaskTags()
|
||||
// 迁移任务绑定的环境变量到通用的数据关联表中
|
||||
migrateTaskEnvs()
|
||||
return nil
|
||||
}
|
||||
|
||||
// migrateTaskEnvs 迁移旧的任务绑定环境变量到通用数据关联表
|
||||
func migrateTaskEnvs() {
|
||||
// 检查 settings 表中是否已经记录了迁移状态
|
||||
if DB.Migrator().HasTable(&models.Setting{}) {
|
||||
var setting models.Setting
|
||||
res := DB.Where(&models.Setting{Section: constant.SectionSystem, Key: constant.KeyTaskEnvsMigrated}).Limit(1).Find(&setting)
|
||||
if res.RowsAffected > 0 && string(setting.Value) == "true" {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if !DB.Migrator().HasColumn(&models.Task{}, "envs") {
|
||||
markTaskEnvsMigrated()
|
||||
return
|
||||
}
|
||||
logger.Infof("[Database] 正在迁移旧任务环境变量绑定...")
|
||||
|
||||
type TaskMigration struct {
|
||||
ID string
|
||||
Envs models.BigText
|
||||
}
|
||||
var tasks []TaskMigration
|
||||
DB.Table((&models.Task{}).TableName()).Select("id, envs").Where("envs IS NOT NULL AND envs != ?", "").Find(&tasks)
|
||||
|
||||
for _, task := range tasks {
|
||||
envs := strings.Split(string(task.Envs), ",")
|
||||
for _, envID := range envs {
|
||||
envID = strings.TrimSpace(envID)
|
||||
if envID == "" {
|
||||
continue
|
||||
}
|
||||
var count int64
|
||||
DB.Model(&models.DataRelation{}).Where("data_id = ? AND relate_id = ? AND type = ?", task.ID, envID, constant.RelationTypeTaskEnv).Count(&count)
|
||||
if count == 0 {
|
||||
relation := models.DataRelation{
|
||||
ID: xid.New().String(),
|
||||
DataID: task.ID,
|
||||
RelateID: envID,
|
||||
Type: constant.RelationTypeTaskEnv,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
DB.Create(&relation)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// if err := DB.Migrator().DropColumn(&models.Task{}, "envs"); err != nil {
|
||||
// logger.Debugf("[Database] 移除 bh_tasks.envs 字段失败: %v", err)
|
||||
// } else {
|
||||
// logger.Infof("[Database] 成功迁移 %d 个环境变量绑定的任务,并删除了旧 envs 字段", len(tasks))
|
||||
// }
|
||||
logger.Infof("[Database] 成功迁移 %d 个环境变量绑定的任务", len(tasks))
|
||||
markTaskEnvsMigrated()
|
||||
}
|
||||
|
||||
func markTaskEnvsMigrated() {
|
||||
if !DB.Migrator().HasTable(&models.Setting{}) {
|
||||
return
|
||||
}
|
||||
var setting models.Setting
|
||||
res := DB.Where(&models.Setting{Section: constant.SectionSystem, Key: constant.KeyTaskEnvsMigrated}).Limit(1).Find(&setting)
|
||||
if res.RowsAffected > 0 {
|
||||
DB.Model(&setting).Update("value", models.BigText("true"))
|
||||
} else {
|
||||
DB.Create(&models.Setting{
|
||||
ID: xid.New().String(),
|
||||
Section: constant.SectionSystem,
|
||||
Key: constant.KeyTaskEnvsMigrated,
|
||||
Value: models.BigText("true"),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// migrateTaskTags 迁移旧的任务标签到通用数据关联表
|
||||
func migrateTaskTags() {
|
||||
// 检查 settings 表中是否已经记录了迁移状态
|
||||
if DB.Migrator().HasTable(&models.Setting{}) {
|
||||
var setting models.Setting
|
||||
res := DB.Where(&models.Setting{Section: constant.SectionSystem, Key: constant.KeyTaskTagsMigrated}).Limit(1).Find(&setting)
|
||||
if res.RowsAffected > 0 && string(setting.Value) == "true" {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if !DB.Migrator().HasColumn(&models.Task{}, "tags") {
|
||||
markTaskTagsMigrated()
|
||||
return
|
||||
}
|
||||
logger.Infof("[Database] 正在迁移旧任务标签...")
|
||||
|
||||
type TaskMigration struct {
|
||||
ID string
|
||||
Tags string
|
||||
}
|
||||
var tasks []TaskMigration
|
||||
DB.Table((&models.Task{}).TableName()).Select("id, tags").Where("tags != ?", "").Find(&tasks)
|
||||
|
||||
for _, task := range tasks {
|
||||
tags := strings.Split(task.Tags, ",")
|
||||
for _, tag := range tags {
|
||||
tag = strings.TrimSpace(tag)
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
var storage models.DataStorage
|
||||
res := DB.Where("type = ? AND name = ?", constant.RelationTypeTaskTag, tag).Limit(1).Find(&storage)
|
||||
if res.RowsAffected == 0 {
|
||||
storage = models.DataStorage{
|
||||
ID: xid.New().String(),
|
||||
Type: constant.RelationTypeTaskTag,
|
||||
Name: tag,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
DB.Create(&storage)
|
||||
}
|
||||
var count int64
|
||||
DB.Model(&models.DataRelation{}).Where("data_id = ? AND relate_id = ? AND type = ?", task.ID, storage.ID, constant.RelationTypeTaskTag).Count(&count)
|
||||
if count == 0 {
|
||||
relation := models.DataRelation{
|
||||
ID: xid.New().String(),
|
||||
DataID: task.ID,
|
||||
RelateID: storage.ID,
|
||||
Type: constant.RelationTypeTaskTag,
|
||||
CreatedAt: models.Now(),
|
||||
UpdatedAt: models.Now(),
|
||||
}
|
||||
DB.Create(&relation)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// if err := DB.Migrator().DropColumn(&models.Task{}, "tags"); err != nil {
|
||||
// logger.Debugf("[Database] 移除 bh_tasks.tags 字段失败: %v", err)
|
||||
// } else {
|
||||
// logger.Infof("[Database] 成功迁移 %d 个带有标签的任务,并删除了旧 tags 字段", len(tasks))
|
||||
// }
|
||||
logger.Infof("[Database] 成功迁移 %d 个带有标签的任务", len(tasks))
|
||||
markTaskTagsMigrated()
|
||||
}
|
||||
|
||||
func markTaskTagsMigrated() {
|
||||
if !DB.Migrator().HasTable(&models.Setting{}) {
|
||||
return
|
||||
}
|
||||
var setting models.Setting
|
||||
res := DB.Where(&models.Setting{Section: constant.SectionSystem, Key: constant.KeyTaskTagsMigrated}).Limit(1).Find(&setting)
|
||||
if res.RowsAffected > 0 {
|
||||
DB.Model(&setting).Update("value", models.BigText("true"))
|
||||
} else {
|
||||
DB.Create(&models.Setting{
|
||||
ID: xid.New().String(),
|
||||
Section: constant.SectionSystem,
|
||||
Key: constant.KeyTaskTagsMigrated,
|
||||
Value: models.BigText("true"),
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package eventbus
|
||||
|
||||
import "sync"
|
||||
|
||||
// Event 统一定义的事件体内包含的数据
|
||||
type Event struct {
|
||||
Type string
|
||||
Payload interface{}
|
||||
}
|
||||
|
||||
// Handler 事件具体的执行句柄
|
||||
type Handler func(event Event)
|
||||
|
||||
// Subscriber 事件订阅者接口,各业务 Service 若关注系统总线可实现此接口
|
||||
type Subscriber interface {
|
||||
SubscribeEvents(bus *EventBus)
|
||||
}
|
||||
|
||||
type EventBus struct {
|
||||
handlers map[string][]Handler
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
func New() *EventBus {
|
||||
return &EventBus{
|
||||
handlers: make(map[string][]Handler),
|
||||
}
|
||||
}
|
||||
|
||||
// Subscribe 注册订阅事件
|
||||
func (bus *EventBus) Subscribe(eventType string, handler Handler) {
|
||||
bus.mu.Lock()
|
||||
defer bus.mu.Unlock()
|
||||
bus.handlers[eventType] = append(bus.handlers[eventType], handler)
|
||||
}
|
||||
|
||||
// Publish 异步抛出事件
|
||||
func (bus *EventBus) Publish(event Event) {
|
||||
bus.mu.RLock()
|
||||
handlers := bus.handlers[event.Type]
|
||||
bus.mu.RUnlock()
|
||||
|
||||
for _, handler := range handlers {
|
||||
// 采用 Goroutine 异步不阻塞核心主线
|
||||
go handler(event)
|
||||
}
|
||||
}
|
||||
|
||||
// 全局唯一的事件总线实例
|
||||
var DefaultBus = New()
|
||||
@@ -0,0 +1,239 @@
|
||||
package executor
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/systime"
|
||||
|
||||
"github.com/robfig/cron/v3"
|
||||
)
|
||||
|
||||
// 东八区时区(默认)
|
||||
var defaultLocation = systime.CST
|
||||
|
||||
// CronManager 统一的任务调度管理器
|
||||
type CronManager struct {
|
||||
cron *cron.Cron
|
||||
scheduler *Scheduler
|
||||
entryMap map[string]cron.EntryID // task ID -> cron entry ID
|
||||
mu sync.RWMutex
|
||||
logger SchedulerLogger
|
||||
OnTrigger func(task CronTask) *ExecutionRequest // 任务触发时的请求构造工厂
|
||||
}
|
||||
|
||||
// NewCronManager 创建一个新的计划任务管理器
|
||||
func NewCronManager(scheduler *Scheduler) *CronManager {
|
||||
// 使用秒级精度的 cron parser
|
||||
c := cron.New(cron.WithSeconds(), cron.WithLocation(defaultLocation))
|
||||
|
||||
m := &CronManager{
|
||||
cron: c,
|
||||
scheduler: scheduler,
|
||||
entryMap: make(map[string]cron.EntryID),
|
||||
logger: &DefaultLogger{},
|
||||
}
|
||||
|
||||
if scheduler != nil && scheduler.logger != nil {
|
||||
m.logger = scheduler.logger
|
||||
}
|
||||
|
||||
return m
|
||||
}
|
||||
|
||||
// SetLogger 设置自定义日志实现
|
||||
func (m *CronManager) SetLogger(logger SchedulerLogger) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.logger = logger
|
||||
}
|
||||
|
||||
// SetScheduler 更新关联的调度器实例
|
||||
func (m *CronManager) SetScheduler(scheduler *Scheduler) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.scheduler = scheduler
|
||||
}
|
||||
|
||||
// Start 启动调度器
|
||||
func (m *CronManager) Start() {
|
||||
m.cron.Start()
|
||||
m.logger.Infof("[CronManager] 调度管理服务已启动")
|
||||
}
|
||||
|
||||
// Stop 停止调度器
|
||||
func (m *CronManager) Stop() {
|
||||
ctx := m.cron.Stop()
|
||||
<-ctx.Done()
|
||||
m.logger.Infof("[CronManager] 调度管理服务已停止")
|
||||
}
|
||||
|
||||
// AddTask 添加或更新计划任务
|
||||
func (m *CronManager) AddTask(task CronTask) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
taskID := task.GetID()
|
||||
|
||||
// 如果已存在,先移除旧的
|
||||
if entryID, exists := m.entryMap[taskID]; exists {
|
||||
m.cron.Remove(entryID)
|
||||
delete(m.entryMap, taskID)
|
||||
}
|
||||
|
||||
// 准备任务执行函数
|
||||
cmd := task.GetCommand()
|
||||
name := task.GetName()
|
||||
timeout := task.GetTimeout()
|
||||
workDir := task.GetWorkDir()
|
||||
envs := task.GetEnvs()
|
||||
languages := task.GetLanguages()
|
||||
useMise := task.UseMise()
|
||||
secrets := task.GetSecrets()
|
||||
|
||||
schedule := strings.TrimSpace(task.GetSchedule())
|
||||
entryID, err := m.cron.AddFunc(schedule, func() {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
m.logger.Errorf("[CronManager] 任务 #%s 执行过程中发生 Panic: %v", taskID, r)
|
||||
}
|
||||
}()
|
||||
|
||||
// 构造执行请求的 Builder
|
||||
reqBuilder := func() *ExecutionRequest {
|
||||
if m.OnTrigger != nil {
|
||||
return m.OnTrigger(task)
|
||||
}
|
||||
return &ExecutionRequest{
|
||||
TaskID: taskID,
|
||||
Name: name,
|
||||
Command: cmd,
|
||||
PreCommand: task.GetPreCommand(),
|
||||
PostCommand: task.GetPostCommand(),
|
||||
Type: TaskTypeCron,
|
||||
Timeout: timeout,
|
||||
WorkDir: workDir,
|
||||
Envs: func() []string {
|
||||
if vars := task.GetEnvVars(); len(vars) > 0 {
|
||||
return vars
|
||||
}
|
||||
return ParseEnvVars(envs)
|
||||
}(),
|
||||
Secrets: secrets,
|
||||
Languages: languages,
|
||||
UseMise: useMise,
|
||||
}
|
||||
}
|
||||
|
||||
randomRange := task.GetRandomRange()
|
||||
if randomRange > 0 && m.scheduler != nil {
|
||||
// 生成 0 到 randomRange 之间的随机秒数
|
||||
delaySeconds := rand.Intn(randomRange)
|
||||
delay := time.Duration(delaySeconds) * time.Second
|
||||
m.logger.Infof("[CronManager] 任务 %s (#%s) 将随机延迟 %v (范围: %ds) 后入队", name, taskID, delay, randomRange)
|
||||
|
||||
// 使用调度器的延时投递功能,不阻塞当前 Cron 协程
|
||||
m.scheduler.EnqueueDelayed(delay, reqBuilder)
|
||||
} else {
|
||||
m.logger.Infof("[CronManager] 触发计划任务: %s (#%s)", name, taskID)
|
||||
if m.scheduler != nil {
|
||||
m.scheduler.EnqueueOrExecute(reqBuilder())
|
||||
}
|
||||
}
|
||||
|
||||
// 触发下次运行时间更新事件
|
||||
m.triggerNextRunEvent(taskID, &ExecutionRequest{TaskID: taskID})
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
m.logger.Errorf("[CronManager] 添加任务失败 #%s: %v", taskID, err)
|
||||
return err
|
||||
}
|
||||
|
||||
m.entryMap[taskID] = entryID
|
||||
m.logger.Infof("[CronManager] 已添加调度: %s (#%s) [%s]", name, taskID, task.GetSchedule())
|
||||
|
||||
// 初始触发一次下次运行时间通知
|
||||
go func() {
|
||||
req := &ExecutionRequest{
|
||||
TaskID: taskID,
|
||||
Name: name,
|
||||
Type: TaskTypeCron,
|
||||
UseMise: task.UseMise(),
|
||||
}
|
||||
m.triggerNextRunEvent(taskID, req)
|
||||
}()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveTask 移除计划任务
|
||||
func (m *CronManager) RemoveTask(taskID string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if entryID, exists := m.entryMap[taskID]; exists {
|
||||
m.cron.Remove(entryID)
|
||||
delete(m.entryMap, taskID)
|
||||
m.logger.Infof("[CronManager] 任务已移除 #%s", taskID)
|
||||
}
|
||||
}
|
||||
|
||||
// triggerNextRunEvent 触发下次运行时间更新事件
|
||||
func (m *CronManager) triggerNextRunEvent(taskID string, req *ExecutionRequest) {
|
||||
m.mu.RLock()
|
||||
entryID, exists := m.entryMap[taskID]
|
||||
m.mu.RUnlock()
|
||||
|
||||
if !exists {
|
||||
return
|
||||
}
|
||||
|
||||
entry := m.cron.Entry(entryID)
|
||||
if !entry.Next.IsZero() && m.scheduler != nil && m.scheduler.handler != nil {
|
||||
m.scheduler.handler.OnCronNextRun(req, entry.Next)
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateCron 校验 Cron 表达式
|
||||
func (m *CronManager) ValidateCron(expression string) error {
|
||||
expression = strings.TrimSpace(expression)
|
||||
if expression == "" {
|
||||
return fmt.Errorf("cron 表达式不能为空")
|
||||
}
|
||||
|
||||
// 如果不是以 @ 开头的描述符,检查位数
|
||||
if !strings.HasPrefix(expression, "@") {
|
||||
fields := strings.Fields(expression)
|
||||
if len(fields) != 6 {
|
||||
return fmt.Errorf("cron 表达式必须为 6 位 (秒 分 时 日 月 周)")
|
||||
}
|
||||
}
|
||||
|
||||
parser := cron.NewParser(cron.Second | cron.Minute | cron.Hour | cron.Dom | cron.Month | cron.Dow | cron.Descriptor)
|
||||
_, err := parser.Parse(expression)
|
||||
return err
|
||||
}
|
||||
|
||||
// GetEntry 获取任务详情
|
||||
func (m *CronManager) GetEntry(taskID string) (cron.Entry, bool) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
|
||||
entryID, exists := m.entryMap[taskID]
|
||||
if !exists {
|
||||
return cron.Entry{}, false
|
||||
}
|
||||
|
||||
return m.cron.Entry(entryID), true
|
||||
}
|
||||
|
||||
// GetScheduledCount 获取已调度任务总数
|
||||
func (m *CronManager) GetScheduledCount() int {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return len(m.entryMap)
|
||||
}
|
||||
@@ -0,0 +1,388 @@
|
||||
package executor
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/creack/pty"
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
// Task 任务基础接口
|
||||
type Task interface {
|
||||
GetID() string
|
||||
GetName() string
|
||||
GetCommand() string
|
||||
GetPreCommand() string
|
||||
GetPostCommand() string
|
||||
GetTimeout() int
|
||||
GetWorkDir() string
|
||||
GetEnvs() string
|
||||
GetEnvVars() []string
|
||||
GetLanguages() []map[string]string
|
||||
GetUseMise() bool
|
||||
}
|
||||
|
||||
// CronTask 计划任务接口
|
||||
type CronTask interface {
|
||||
Task
|
||||
GetSchedule() string
|
||||
UseMise() bool
|
||||
GetSecrets() []string
|
||||
GetRandomRange() int
|
||||
}
|
||||
|
||||
// Request 任务执行请求
|
||||
type Request struct {
|
||||
Command string
|
||||
PreCommand string
|
||||
PostCommand string
|
||||
WorkDir string
|
||||
Envs []string
|
||||
Timeout int // 任务超时时间(分钟)
|
||||
Languages []map[string]string
|
||||
UseMise bool
|
||||
}
|
||||
|
||||
// Result 任务执行结果
|
||||
type Result struct {
|
||||
Output string
|
||||
Error string
|
||||
Status string // 状态: success, failed
|
||||
Duration int64 // 毫秒
|
||||
ExitCode int
|
||||
StartTime time.Time
|
||||
EndTime time.Time
|
||||
}
|
||||
|
||||
// Hooks 执行钩子接口
|
||||
type Hooks interface {
|
||||
// PreExecute 执行前钩子,返回日志ID和错误
|
||||
PreExecute(ctx context.Context, req Request) (logID string, err error)
|
||||
|
||||
// PostExecute 执行后钩子,处理日志压缩和记录更新
|
||||
PostExecute(ctx context.Context, logID string, result *Result) error
|
||||
|
||||
// OnHeartbeat 执行中心跳钩子,用于更新实时状态
|
||||
OnHeartbeat(ctx context.Context, logID string, duration int64) error
|
||||
}
|
||||
|
||||
// Execute 执行命令(基础版本,不带钩子)
|
||||
func Execute(ctx context.Context, req Request, stdout, stderr io.Writer) (*Result, error) {
|
||||
return ExecuteWithHooks(ctx, req, stdout, stderr, nil)
|
||||
}
|
||||
|
||||
// ExecuteWithHooks 执行命令(带钩子支持)
|
||||
func ExecuteWithHooks(ctx context.Context, req Request, stdout, stderr io.Writer, hooks Hooks) (*Result, error) {
|
||||
start := time.Now()
|
||||
|
||||
// 演示模式拦截
|
||||
if constant.DemoMode {
|
||||
logger.Warnf("[Executor] 演示模式下已拦截命令执行: %s", req.Command)
|
||||
if stdout != nil {
|
||||
stdout.Write([]byte("\r\n\033[1;33m[演示模式] 命令执行已跳过\033[0m\r\n"))
|
||||
}
|
||||
|
||||
// 仍然触发 PreExecute 以便流程完整
|
||||
var logID string
|
||||
if hooks != nil {
|
||||
logID, _ = hooks.PreExecute(ctx, req)
|
||||
}
|
||||
|
||||
result := &Result{
|
||||
Status: constant.TaskStatusFailed,
|
||||
Output: "[演示模式] 该任务在演示模式下被禁用执行",
|
||||
StartTime: start,
|
||||
EndTime: time.Now(),
|
||||
}
|
||||
|
||||
if hooks != nil {
|
||||
hooks.PostExecute(ctx, logID, result)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// 2. 执行命令
|
||||
timeout := req.Timeout
|
||||
var execCtx context.Context
|
||||
var cancel context.CancelFunc
|
||||
|
||||
if timeout > 0 {
|
||||
execCtx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Minute)
|
||||
} else {
|
||||
execCtx, cancel = context.WithCancel(ctx)
|
||||
}
|
||||
defer cancel()
|
||||
|
||||
// 如果指定使用 mise,则预先构建好带 mise 的命令,这样 PreExecute 记录的就是完整命令
|
||||
if req.UseMise {
|
||||
utils.InjectNodePath(&req.Envs, req.Languages)
|
||||
req.Command = utils.BuildMiseCommand(req.Command, req.Languages)
|
||||
req.UseMise = false
|
||||
}
|
||||
|
||||
// 组合指令(如果存在前置或后置指令)
|
||||
if req.PreCommand != "" || req.PostCommand != "" {
|
||||
finalCmd := ""
|
||||
if req.PreCommand != "" {
|
||||
finalCmd += req.PreCommand + "\n"
|
||||
}
|
||||
finalCmd += req.Command
|
||||
if req.PostCommand != "" {
|
||||
finalCmd += "\n" + req.PostCommand
|
||||
}
|
||||
req.Command = finalCmd
|
||||
}
|
||||
|
||||
// 1. 执行前钩子
|
||||
var logID string
|
||||
if hooks != nil {
|
||||
id, err := hooks.PreExecute(ctx, req)
|
||||
if err != nil {
|
||||
return &Result{
|
||||
Status: constant.TaskStatusFailed,
|
||||
Duration: 0,
|
||||
ExitCode: 1,
|
||||
StartTime: start,
|
||||
EndTime: time.Now(),
|
||||
}, err
|
||||
}
|
||||
logID = id
|
||||
}
|
||||
|
||||
shell, args := utils.GetShellCommand(req.Command)
|
||||
cmd := exec.CommandContext(execCtx, shell, args...)
|
||||
|
||||
usePty := runtime.GOOS != "windows" && stdout != nil && (stdout == stderr || stdout == io.Discard)
|
||||
SetProcessGroupAndCancel(cmd, usePty)
|
||||
|
||||
// 设置工作目录
|
||||
// 设置工作目录
|
||||
workDir := strings.TrimSpace(req.WorkDir)
|
||||
if workDir != "" {
|
||||
cmd.Dir = workDir
|
||||
}
|
||||
|
||||
// 设置环境变量(始终继承系统环境变量)
|
||||
cmd.Env = os.Environ()
|
||||
if len(req.Envs) > 0 {
|
||||
cmd.Env = append(cmd.Env, req.Envs...)
|
||||
}
|
||||
// 强制注入终端环境标识及禁用输出缓冲的标志
|
||||
cmd.Env = append(cmd.Env,
|
||||
"TERM=xterm",
|
||||
"PYTHONUNBUFFERED=1",
|
||||
"NODE_NO_WARNINGS=1",
|
||||
)
|
||||
|
||||
var pipeWriter *os.File
|
||||
var ptyFile *os.File
|
||||
var copyDone chan struct{}
|
||||
var err error
|
||||
|
||||
var started bool
|
||||
// 尝试开启 PTY 模式(Unix/macOS 且输出合并时)
|
||||
if runtime.GOOS != "windows" && stdout != nil && (stdout == stderr || stdout == io.Discard) {
|
||||
// 强制注入终端环境标识及禁用输出缓冲的标志,确保 PTY 模式下最佳实时性能
|
||||
cmd.Env = append(cmd.Env,
|
||||
"TERM=xterm",
|
||||
"PYTHONUNBUFFERED=1",
|
||||
"NODE_NO_WARNINGS=1",
|
||||
)
|
||||
f, ptyErr := pty.Start(cmd)
|
||||
if ptyErr == nil {
|
||||
logger.Infof("[Executor] #%s 启动于 PTY 模式", logID)
|
||||
ptyFile = f
|
||||
started = true
|
||||
copyDone = make(chan struct{})
|
||||
go func() {
|
||||
defer close(copyDone)
|
||||
// io.Copy 对于 PTY 来说是最稳健且即时的流式拷贝
|
||||
io.Copy(stdout, f)
|
||||
f.Close()
|
||||
}()
|
||||
} else {
|
||||
logger.Errorf("[Executor] 任务 #%s PTY 启动失败: %v", logID, ptyErr)
|
||||
}
|
||||
}
|
||||
|
||||
if !started {
|
||||
// 如果 stdout 和 stderr 指针不一致,但在逻辑上我们知道它们是同一个 MultiWriter,
|
||||
// 这里会显示为 Pipe 模式。
|
||||
if stdout != stderr && stdout != io.Discard {
|
||||
logger.Debugf("[Executor] 任务 #%d stdout (%p) 和 stderr (%p) 不同,回退到 Pipe 模式。", logID, stdout, stderr)
|
||||
}
|
||||
logger.Infof("[Executor] #%s 启动于 Pipe 模式", logID)
|
||||
if stdout != nil && stdout == stderr {
|
||||
pr, pw, err := os.Pipe()
|
||||
if err == nil {
|
||||
cmd.Stdout = pw
|
||||
cmd.Stderr = pw
|
||||
pipeWriter = pw
|
||||
copyDone = make(chan struct{})
|
||||
go func() {
|
||||
io.Copy(stdout, pr)
|
||||
pr.Close()
|
||||
close(copyDone)
|
||||
}()
|
||||
} else {
|
||||
cmd.Stdout = stdout
|
||||
cmd.Stderr = stderr
|
||||
}
|
||||
} else {
|
||||
cmd.Stdout = stdout
|
||||
cmd.Stderr = stderr
|
||||
}
|
||||
|
||||
// 使用 cmd.Start() + Wait() 以便在后台处理心跳
|
||||
err = cmd.Start()
|
||||
if err != nil {
|
||||
if pipeWriter != nil {
|
||||
pipeWriter.Close()
|
||||
}
|
||||
// 启动失败的处理
|
||||
end := time.Now()
|
||||
result := &Result{
|
||||
Status: constant.TaskStatusFailed,
|
||||
Duration: end.Sub(start).Milliseconds(),
|
||||
ExitCode: 1,
|
||||
StartTime: start, // 记录开始时间
|
||||
EndTime: end,
|
||||
}
|
||||
// 执行后钩子
|
||||
if hooks != nil {
|
||||
result.Output += "\n[系统错误] " + err.Error()
|
||||
hooks.PostExecute(ctx, logID, result)
|
||||
}
|
||||
return result, err
|
||||
}
|
||||
|
||||
// 在父进程中关闭写端,这样子进程退出后 pr 才会收到 EOF
|
||||
if pipeWriter != nil {
|
||||
pipeWriter.Close()
|
||||
}
|
||||
} else {
|
||||
// PTY 模式下 cmd.Start() 已经在 pty.Start(cmd) 中调用过了
|
||||
}
|
||||
|
||||
// 启动心跳协程
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
// 每3秒一次心跳
|
||||
ticker := time.NewTicker(3 * time.Second)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-done:
|
||||
return
|
||||
case <-ticker.C:
|
||||
if hooks != nil {
|
||||
hooks.OnHeartbeat(ctx, logID, time.Since(start).Milliseconds())
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
// 等待命令完成
|
||||
err = cmd.Wait()
|
||||
close(done) // 停止心跳
|
||||
|
||||
// PTY 模式下需要显式关闭
|
||||
if ptyFile != nil {
|
||||
ptyFile.Close()
|
||||
}
|
||||
|
||||
// 等待日志复制完成
|
||||
if copyDone != nil {
|
||||
<-copyDone
|
||||
}
|
||||
|
||||
end := time.Now()
|
||||
|
||||
result := &Result{
|
||||
StartTime: start,
|
||||
EndTime: end,
|
||||
Duration: end.Sub(start).Milliseconds(),
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
result.Status = constant.TaskStatusFailed
|
||||
result.Error = err.Error()
|
||||
if exitErr, ok := err.(*exec.ExitError); ok {
|
||||
result.ExitCode = exitErr.ExitCode()
|
||||
} else {
|
||||
result.ExitCode = 1
|
||||
}
|
||||
} else {
|
||||
result.Status = constant.TaskStatusSuccess
|
||||
result.ExitCode = 0
|
||||
}
|
||||
|
||||
// 3. 执行后钩子
|
||||
if hooks != nil {
|
||||
if hookErr := hooks.PostExecute(ctx, logID, result); hookErr != nil {
|
||||
// 记录钩子错误但不影响执行结果
|
||||
result.Output += "\n[钩子错误] " + hookErr.Error()
|
||||
}
|
||||
}
|
||||
|
||||
return result, err
|
||||
}
|
||||
|
||||
// ParseEnvVars 解析环境变量字符串 "KEY1=VALUE1,KEY2=VALUE2"
|
||||
func ParseEnvVars(envStr string) []string {
|
||||
if envStr == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
pairs := strings.Split(envStr, ",")
|
||||
result := make([]string, 0, len(pairs))
|
||||
|
||||
for _, pair := range pairs {
|
||||
if pair == "" {
|
||||
continue
|
||||
}
|
||||
// 解码特殊字符
|
||||
pair = strings.ReplaceAll(pair, "{{COMMA}}", ",")
|
||||
pair = strings.ReplaceAll(pair, "{{EQUAL}}", "=")
|
||||
pair = strings.ReplaceAll(pair, "{{NL}}", "\n")
|
||||
result = append(result, pair)
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// FormatEnvVars 将环境变量列表格式化为逗号分隔的字符串 "KEY1=VALUE1,KEY2=VALUE2"
|
||||
// 会对 , 和 = 以及换行符进行转义
|
||||
func FormatEnvVars(envs []string) string {
|
||||
if len(envs) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
pairs := make([]string, 0, len(envs))
|
||||
for _, pair := range envs {
|
||||
// 寻找第一个等号
|
||||
idx := strings.Index(pair, "=")
|
||||
if idx == -1 {
|
||||
continue
|
||||
}
|
||||
name := pair[:idx]
|
||||
value := pair[idx+1:]
|
||||
|
||||
// 转义特殊字符
|
||||
encodedValue := strings.ReplaceAll(value, ",", "{{COMMA}}")
|
||||
encodedValue = strings.ReplaceAll(encodedValue, "=", "{{EQUAL}}")
|
||||
encodedValue = strings.ReplaceAll(encodedValue, "\n", "{{NL}}")
|
||||
pairs = append(pairs, fmt.Sprintf("%s=%s", name, encodedValue))
|
||||
}
|
||||
|
||||
return strings.Join(pairs, ",")
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
//go:build !windows
|
||||
|
||||
package executor
|
||||
|
||||
import (
|
||||
"os/exec"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
func SetProcessGroupAndCancel(cmd *exec.Cmd, usePty bool) {
|
||||
if !usePty {
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
||||
}
|
||||
cmd.Cancel = func() error {
|
||||
if cmd.Process != nil {
|
||||
// Kill the entire process group by sending SIGKILL to negative PID
|
||||
return syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
//go:build windows
|
||||
|
||||
package executor
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os/exec"
|
||||
)
|
||||
|
||||
func SetProcessGroupAndCancel(cmd *exec.Cmd, usePty bool) {
|
||||
cmd.Cancel = func() error {
|
||||
if cmd.Process != nil {
|
||||
killCmd := exec.Command("taskkill", "/F", "/T", "/PID", fmt.Sprintf("%d", cmd.Process.Pid))
|
||||
return killCmd.Run()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,690 @@
|
||||
package executor
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
// safeBuffer 一个线程安全的字节缓冲区,用于合并 stdout 和 stderr
|
||||
type safeBuffer struct {
|
||||
mu sync.Mutex
|
||||
buf bytes.Buffer
|
||||
}
|
||||
|
||||
func (s *safeBuffer) Write(p []byte) (n int, err error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.buf.Write(p)
|
||||
}
|
||||
|
||||
func (s *safeBuffer) String() string {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.buf.String()
|
||||
}
|
||||
|
||||
// SchedulerConfig 调度器配置
|
||||
type SchedulerConfig struct {
|
||||
WorkerCount int // Worker 数量
|
||||
QueueSize int // 队列大小
|
||||
RateInterval time.Duration // 速率限制间隔
|
||||
Verbose bool // 是否开启详细日志
|
||||
StrictQueue bool // 是否开启严格排队(满时拒绝执行,不降级直接执行)
|
||||
}
|
||||
|
||||
// TaskType 任务类型
|
||||
type TaskType string
|
||||
|
||||
const (
|
||||
TaskTypeCron TaskType = "cron" // 计划任务
|
||||
TaskTypeManual TaskType = "manual" // 手动任务
|
||||
TaskTypeSystem TaskType = "system" // 系统任务
|
||||
)
|
||||
|
||||
// TaskStatus 任务状态
|
||||
type TaskStatus string
|
||||
|
||||
const (
|
||||
TaskStatusPending TaskStatus = TaskStatus(constant.TaskStatusPending) // 等待中
|
||||
TaskStatusRunning TaskStatus = TaskStatus(constant.TaskStatusRunning) // 运行中
|
||||
TaskStatusSuccess TaskStatus = TaskStatus(constant.TaskStatusSuccess) // 成功
|
||||
TaskStatusFailed TaskStatus = TaskStatus(constant.TaskStatusFailed) // 失败
|
||||
TaskStatusTimeout TaskStatus = TaskStatus(constant.TaskStatusTimeout) // 超时
|
||||
TaskStatusCancelled TaskStatus = TaskStatus(constant.TaskStatusCancelled) // 已取消
|
||||
)
|
||||
|
||||
// ExecutionRequest 执行请求(标准接口)
|
||||
type ExecutionRequest struct {
|
||||
TaskID string // 任务 ID
|
||||
LogID string // 日志 ID
|
||||
Name string // 任务名称
|
||||
Type TaskType // 任务类型
|
||||
Command string // 命令
|
||||
MaskedCommand string // 脱敏后的命令(用于日志和展示)
|
||||
PreCommand string // 前置命令
|
||||
PostCommand string // 后置命令
|
||||
WorkDir string // 工作目录
|
||||
Envs []string // 环境变量
|
||||
Secrets []string // 需要脱敏的密码
|
||||
Timeout int // 超时时间(分钟)
|
||||
Languages []map[string]string // 语言环境配置
|
||||
UseMise bool // 是否使用 mise
|
||||
Metadata ExecutionMetadata // 额外元数据
|
||||
}
|
||||
|
||||
// ExecutionMetadata 执行额外元数据
|
||||
type ExecutionMetadata struct {
|
||||
GoID int64 // 关联的 goroutine ID
|
||||
RetryIndex int // 当前重试索引
|
||||
}
|
||||
|
||||
// ExecutionResult 执行结果(标准接口)
|
||||
type ExecutionResult struct {
|
||||
TaskID string // 任务 ID
|
||||
LogID string // 日志 ID
|
||||
Success bool // 是否成功
|
||||
Output string // 输出内容
|
||||
Error string // 错误信息
|
||||
Status string // 状态: success, failed, timeout, cancelled
|
||||
Duration int64 // 执行时长(毫秒)
|
||||
ExitCode int // 退出码
|
||||
StartTime time.Time // 开始时间
|
||||
EndTime time.Time // 结束时间
|
||||
}
|
||||
|
||||
// SchedulerEventHandler 调度器事件处理器(标准接口)
|
||||
// 主服务端和 Agent 端通过实现不同的 Handler 来处理事件
|
||||
type SchedulerEventHandler interface {
|
||||
// OnTaskScheduled 任务被调度(加入队列)时触发
|
||||
OnTaskScheduled(req *ExecutionRequest)
|
||||
|
||||
// OnTaskExecuting 任务准备开始执行时触发
|
||||
// 返回 stdout/stderr 写入器用于实时日志推送
|
||||
// 主服务端:返回 TinyLog 写入器(写入本地文件)
|
||||
// Agent 端:返回 WebSocket 写入器(实时推送到主服务)
|
||||
OnTaskExecuting(req *ExecutionRequest) (stdout, stderr io.Writer, err error)
|
||||
|
||||
// OnTaskStarted 任务实际开始运行(已经过了队列等待和速率限制)
|
||||
OnTaskStarted(req *ExecutionRequest)
|
||||
|
||||
// OnTaskCompleted 任务执行完成时触发
|
||||
// 主服务端:压缩日志、更新数据库、清理旧日志
|
||||
// Agent 端:通过 WebSocket 发送执行结果到主服务
|
||||
OnTaskCompleted(req *ExecutionRequest, result *ExecutionResult)
|
||||
|
||||
// OnTaskFailed 任务执行失败时触发
|
||||
OnTaskFailed(req *ExecutionRequest, err error)
|
||||
|
||||
// OnCronNextRun 计划任务下次运行时间更新时触发
|
||||
OnCronNextRun(req *ExecutionRequest, nextRun time.Time)
|
||||
|
||||
// OnTaskHeartbeat 任务执行心跳(用于更新实时耗时等)
|
||||
OnTaskHeartbeat(req *ExecutionRequest, duration int64)
|
||||
}
|
||||
|
||||
// SchedulerLogger 日志接口(允许自定义日志实现)
|
||||
type SchedulerLogger interface {
|
||||
Infof(format string, args ...interface{})
|
||||
Warnf(format string, args ...interface{})
|
||||
Errorf(format string, args ...interface{})
|
||||
}
|
||||
|
||||
// DefaultLogger 默认日志实现(使用 fmt)
|
||||
type DefaultLogger struct{}
|
||||
|
||||
func (l *DefaultLogger) Infof(format string, args ...interface{}) {
|
||||
fmt.Printf("[INFO] "+format+"\n", args...)
|
||||
}
|
||||
func (l *DefaultLogger) Warnf(format string, args ...interface{}) {
|
||||
fmt.Printf("[WARN] "+format+"\n", args...)
|
||||
}
|
||||
func (l *DefaultLogger) Errorf(format string, args ...interface{}) {
|
||||
fmt.Printf("[ERROR] "+format+"\n", args...)
|
||||
}
|
||||
|
||||
// schedulerHooksAdapter 适配器:将 executor.Hooks 映射到 SchedulerEventHandler
|
||||
type schedulerHooksAdapter struct {
|
||||
handler SchedulerEventHandler
|
||||
req *ExecutionRequest
|
||||
}
|
||||
|
||||
func (h *schedulerHooksAdapter) PreExecute(ctx context.Context, req Request) (string, error) {
|
||||
return h.req.LogID, nil
|
||||
}
|
||||
|
||||
func (h *schedulerHooksAdapter) PostExecute(ctx context.Context, logID string, result *Result) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *schedulerHooksAdapter) OnHeartbeat(ctx context.Context, logID string, duration int64) error {
|
||||
if h.handler != nil {
|
||||
h.handler.OnTaskHeartbeat(h.req, duration)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// TaskExecutor 定义任务执行函数签名
|
||||
type TaskExecutor func(ctx context.Context, req *ExecutionRequest, stdout, stderr io.Writer) (*Result, error)
|
||||
|
||||
// WorkerStatus 定义并发池中单个 Worker 的状态
|
||||
type WorkerStatus struct {
|
||||
ID int `json:"id"`
|
||||
Status string `json:"status"` // 状态: "idle" 或 "running"
|
||||
TaskID string `json:"task_id,omitempty"`
|
||||
TaskName string `json:"task_name,omitempty"`
|
||||
StartTime int64 `json:"start_time,omitempty"` // 开始时间戳 (秒)
|
||||
Duration int64 `json:"duration,omitempty"` // 已运行时长 (秒)
|
||||
}
|
||||
|
||||
// Scheduler 统一调度器(独立组件,可在主服务和 Agent 中复用)
|
||||
// 调度器本身只负责队列管理和任务调度,具体的执行逻辑和事件处理由 Handler 实现
|
||||
type Scheduler struct {
|
||||
config SchedulerConfig
|
||||
handler SchedulerEventHandler
|
||||
executor TaskExecutor
|
||||
taskQueue chan *ExecutionRequest
|
||||
rateLimiter <-chan time.Time
|
||||
stopCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
mu sync.RWMutex
|
||||
logger SchedulerLogger
|
||||
runningTasks map[string]context.CancelFunc // 记录运行中的任务,用于停止 (TaskID -> CancelFunc)
|
||||
runningExecs map[string]context.CancelFunc // 记录运行中的执行,用于停止 (LogID -> CancelFunc)
|
||||
|
||||
workers []WorkerStatus
|
||||
workerMu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewScheduler 创建调度器
|
||||
func NewScheduler(config SchedulerConfig, handler SchedulerEventHandler) *Scheduler {
|
||||
if config.WorkerCount <= 0 {
|
||||
config.WorkerCount = 4
|
||||
}
|
||||
if config.WorkerCount > 1000 {
|
||||
config.WorkerCount = 1000
|
||||
}
|
||||
if config.QueueSize <= 0 {
|
||||
config.QueueSize = 100
|
||||
}
|
||||
if config.QueueSize > 50000 {
|
||||
config.QueueSize = 50000
|
||||
}
|
||||
if config.RateInterval <= 0 {
|
||||
config.RateInterval = 200 * time.Millisecond
|
||||
}
|
||||
|
||||
s := &Scheduler{
|
||||
config: config,
|
||||
handler: handler,
|
||||
executor: func(ctx context.Context, req *ExecutionRequest, stdout, stderr io.Writer) (*Result, error) {
|
||||
hooks := &schedulerHooksAdapter{handler: handler, req: req}
|
||||
return ExecuteWithHooks(ctx, Request{
|
||||
Command: req.Command,
|
||||
PreCommand: req.PreCommand,
|
||||
PostCommand: req.PostCommand,
|
||||
WorkDir: req.WorkDir,
|
||||
Envs: req.Envs,
|
||||
Timeout: req.Timeout,
|
||||
Languages: req.Languages,
|
||||
UseMise: req.UseMise,
|
||||
}, stdout, stderr, hooks)
|
||||
},
|
||||
taskQueue: make(chan *ExecutionRequest, config.QueueSize),
|
||||
rateLimiter: time.Tick(config.RateInterval),
|
||||
stopCh: make(chan struct{}),
|
||||
logger: &DefaultLogger{},
|
||||
runningTasks: make(map[string]context.CancelFunc),
|
||||
runningExecs: make(map[string]context.CancelFunc),
|
||||
workers: make([]WorkerStatus, config.WorkerCount),
|
||||
}
|
||||
|
||||
for i := 0; i < config.WorkerCount; i++ {
|
||||
s.workers[i] = WorkerStatus{
|
||||
ID: i,
|
||||
Status: "idle",
|
||||
}
|
||||
}
|
||||
|
||||
return s
|
||||
}
|
||||
|
||||
// SetLogger 设置自定义日志实现
|
||||
func (s *Scheduler) SetLogger(logger SchedulerLogger) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.logger = logger
|
||||
}
|
||||
|
||||
// SetExecutor 设置任务执行器
|
||||
func (s *Scheduler) SetExecutor(executor TaskExecutor) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.executor = executor
|
||||
}
|
||||
|
||||
// Start 启动调度器
|
||||
func (s *Scheduler) Start() {
|
||||
for i := 0; i < s.config.WorkerCount; i++ {
|
||||
s.wg.Add(1)
|
||||
go s.worker(i)
|
||||
}
|
||||
s.logger.Infof("[Scheduler] 已启动")
|
||||
}
|
||||
|
||||
// Stop 停止调度器
|
||||
func (s *Scheduler) Stop() {
|
||||
close(s.stopCh)
|
||||
s.wg.Wait()
|
||||
s.logger.Infof("[Scheduler] 已停止")
|
||||
}
|
||||
|
||||
// Enqueue 将任务加入队列
|
||||
func (s *Scheduler) Enqueue(req *ExecutionRequest) error {
|
||||
select {
|
||||
case s.taskQueue <- req:
|
||||
if s.handler != nil {
|
||||
s.handler.OnTaskScheduled(req)
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
// 队列满,返回错误
|
||||
return fmt.Errorf("任务队列已满")
|
||||
}
|
||||
}
|
||||
|
||||
// EnqueueOrExecute 将任务加入队列,如果队列满则直接执行
|
||||
func (s *Scheduler) EnqueueOrExecute(req *ExecutionRequest) {
|
||||
select {
|
||||
case s.taskQueue <- req:
|
||||
// 成功入队
|
||||
if s.handler != nil {
|
||||
s.handler.OnTaskScheduled(req)
|
||||
}
|
||||
default:
|
||||
if s.config.StrictQueue {
|
||||
s.logger.Errorf("[Scheduler] 任务队列已满,拒绝执行任务 %s", req.TaskID)
|
||||
if s.handler != nil {
|
||||
s.handler.OnTaskFailed(req, fmt.Errorf("任务队列已满,拒绝执行"))
|
||||
}
|
||||
} else {
|
||||
// 队列满,直接执行(降级处理)
|
||||
s.logger.Warnf("[Scheduler] 任务队列已满,直接执行任务 %s", req.TaskID)
|
||||
go s.executeTask(req)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// EnqueueDelayed 延迟将任务加入队列执行
|
||||
func (s *Scheduler) EnqueueDelayed(delay time.Duration, reqBuilder func() *ExecutionRequest) {
|
||||
go func() {
|
||||
select {
|
||||
case <-time.After(delay):
|
||||
if req := reqBuilder(); req != nil {
|
||||
s.EnqueueOrExecute(req)
|
||||
}
|
||||
case <-s.stopCh:
|
||||
// 调度器停止时取消延迟投递
|
||||
return
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// ExecuteSync 同步执行任务(不经过队列)
|
||||
func (s *Scheduler) ExecuteSync(req *ExecutionRequest) (*ExecutionResult, error) {
|
||||
return s.executeTask(req)
|
||||
}
|
||||
|
||||
// worker 工作协程
|
||||
func (s *Scheduler) worker(id int) {
|
||||
defer s.wg.Done()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-s.stopCh:
|
||||
return
|
||||
case req := <-s.taskQueue:
|
||||
func() {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
s.logger.Errorf("[Scheduler] Worker %d panic while processing task %s: %v", id, req.TaskID, r)
|
||||
}
|
||||
}()
|
||||
// 速率限制
|
||||
<-s.rateLimiter
|
||||
|
||||
func() {
|
||||
// 恢复 worker 状态为空闲
|
||||
defer func() {
|
||||
s.workerMu.Lock()
|
||||
if id >= 0 && id < len(s.workers) {
|
||||
s.workers[id].Status = "idle"
|
||||
s.workers[id].TaskID = ""
|
||||
s.workers[id].TaskName = ""
|
||||
s.workers[id].StartTime = 0
|
||||
}
|
||||
s.workerMu.Unlock()
|
||||
}()
|
||||
|
||||
// 更新 worker 状态为运行中
|
||||
s.workerMu.Lock()
|
||||
if id >= 0 && id < len(s.workers) {
|
||||
s.workers[id].Status = "running"
|
||||
s.workers[id].TaskID = req.TaskID
|
||||
s.workers[id].TaskName = req.Name
|
||||
s.workers[id].StartTime = time.Now().Unix()
|
||||
}
|
||||
s.workerMu.Unlock()
|
||||
|
||||
s.executeTask(req)
|
||||
}()
|
||||
}()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// executeTask 执行任务(本地执行)
|
||||
func (s *Scheduler) executeTask(req *ExecutionRequest) (*ExecutionResult, error) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
s.logger.Errorf("[Scheduler] 任务 %s 执行过程中发生 Panic: %v", req.TaskID, r)
|
||||
}
|
||||
}()
|
||||
start := time.Now()
|
||||
|
||||
s.logger.Infof("[Scheduler] 开始执行: %s (#%s) [%s]", req.Name, req.TaskID, req.Type)
|
||||
|
||||
// 演示模式拦截
|
||||
if constant.DemoMode {
|
||||
s.logger.Infof("[Scheduler] 演示模式下已跳过任务 %s (%s) 的执行", req.TaskID, req.Name)
|
||||
// 仍然触发 OnTaskExecuting 以便创建初始日志记录(业务层面的 Handler 会处理)
|
||||
var stdout io.Writer
|
||||
if s.handler != nil {
|
||||
stdout, _, _ = s.handler.OnTaskExecuting(req)
|
||||
}
|
||||
|
||||
result := &ExecutionResult{
|
||||
TaskID: req.TaskID,
|
||||
LogID: req.LogID,
|
||||
Success: false,
|
||||
Status: constant.TaskStatusFailed,
|
||||
Error: "[演示模式] 该任务在演示模式下被禁用执行",
|
||||
StartTime: start,
|
||||
EndTime: time.Now(),
|
||||
}
|
||||
|
||||
if stdout != nil {
|
||||
stdout.Write([]byte("\r\n\033[1;33m[演示模式] 定时任务/手动任务执行已跳过\033[0m\r\n"))
|
||||
}
|
||||
|
||||
if s.handler != nil {
|
||||
s.handler.OnTaskCompleted(req, result)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// 如果指定使用 mise,则预先构建好带 mise 的命令,这样 OnTaskExecuting 记录的就是完整命令
|
||||
if req.UseMise {
|
||||
// 先注入 NODE_PATH (由于调度器会把 UseMise 置为 false,所以必须在这里提前处理)
|
||||
utils.InjectNodePath(&req.Envs, req.Languages)
|
||||
req.Command = utils.BuildMiseCommand(req.Command, req.Languages)
|
||||
req.UseMise = false
|
||||
}
|
||||
// 确保系统级敏感信息(数据库地址、账号、密码等)始终在脱敏列表中
|
||||
allSecrets := append([]string{}, req.Secrets...)
|
||||
allSecrets = append(allSecrets, utils.GetSystemSecrets()...)
|
||||
|
||||
s.logger.Infof("[Scheduler] 命令: %s", utils.MaskSecrets(req.Command, allSecrets))
|
||||
|
||||
if s.config.Verbose {
|
||||
workDir := req.WorkDir
|
||||
if workDir == "" {
|
||||
workDir, _ = os.Getwd()
|
||||
}
|
||||
s.logger.Infof("[Scheduler] 任务 #%s 进程 UID: %d, GID: %d", req.TaskID, os.Getuid(), os.Getgid())
|
||||
s.logger.Infof("[Scheduler] 任务 #%s 工作目录: %s", req.TaskID, workDir)
|
||||
}
|
||||
|
||||
// 1. 执行前事件:获取 stdout/stderr 写入器
|
||||
var stdout, stderr io.Writer
|
||||
var err error
|
||||
if s.handler != nil {
|
||||
stdout, stderr, err = s.handler.OnTaskExecuting(req)
|
||||
if err != nil {
|
||||
s.logger.Errorf("[Scheduler] 任务 %s 执行前事件失败: %v", req.TaskID, err)
|
||||
if s.handler != nil {
|
||||
s.handler.OnTaskFailed(req, err)
|
||||
}
|
||||
return &ExecutionResult{
|
||||
TaskID: req.TaskID,
|
||||
Success: false,
|
||||
Status: constant.TaskStatusFailed,
|
||||
Error: err.Error(),
|
||||
Duration: 0,
|
||||
ExitCode: 1,
|
||||
StartTime: start,
|
||||
EndTime: time.Now(),
|
||||
}, err
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 准备输出缓冲区(使用合并缓冲区保证顺序)
|
||||
var combinedBuf safeBuffer
|
||||
var stdoutWriter, stderrWriter io.Writer
|
||||
|
||||
if stdout != nil && stdout == stderr {
|
||||
// 如果 stdout 和 stderr 是同一个对象,合并成一个 MultiWriter
|
||||
// 这样后面 ExecuteWithHooks 才能识别出它们是同一个,从而开启 PTY 模式
|
||||
mw := io.MultiWriter(&combinedBuf, stdout)
|
||||
stdoutWriter = mw
|
||||
stderrWriter = mw
|
||||
} else {
|
||||
if stdout != nil {
|
||||
stdoutWriter = io.MultiWriter(&combinedBuf, stdout)
|
||||
} else {
|
||||
stdoutWriter = &combinedBuf
|
||||
}
|
||||
|
||||
if stderr != nil {
|
||||
stderrWriter = io.MultiWriter(&combinedBuf, stderr)
|
||||
} else {
|
||||
stderrWriter = &combinedBuf
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 实际开始执行事件 (经过队列和速率限制之后)
|
||||
if s.handler != nil {
|
||||
s.handler.OnTaskStarted(req)
|
||||
}
|
||||
|
||||
// 4. 执行命令(使用 executor.Execute)
|
||||
// 创建带取消功能的上下文
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
if req.Timeout > 0 {
|
||||
ctx, cancel = context.WithTimeout(ctx, time.Duration(req.Timeout)*time.Minute)
|
||||
}
|
||||
defer cancel()
|
||||
|
||||
// 注册到运行中任务
|
||||
s.mu.Lock()
|
||||
s.runningTasks[req.TaskID] = cancel
|
||||
if req.LogID != "" {
|
||||
s.runningExecs[req.LogID] = cancel
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
defer func() {
|
||||
s.mu.Lock()
|
||||
delete(s.runningTasks, req.TaskID)
|
||||
if req.LogID != "" {
|
||||
delete(s.runningExecs, req.LogID)
|
||||
}
|
||||
s.mu.Unlock()
|
||||
}()
|
||||
|
||||
execResult, execErr := s.executor(ctx, req, stdoutWriter, stderrWriter)
|
||||
|
||||
// 5. 构建结果
|
||||
result := &ExecutionResult{
|
||||
TaskID: req.TaskID,
|
||||
LogID: req.LogID, // 传递 LogID
|
||||
}
|
||||
|
||||
// 统一获取输出并调用封装的脱敏函数
|
||||
rawStr := utils.MaskSecrets(combinedBuf.String(), req.Secrets)
|
||||
|
||||
if execResult != nil {
|
||||
result.Success = execResult.Status == constant.TaskStatusSuccess
|
||||
result.Output = rawStr
|
||||
result.Status = execResult.Status
|
||||
result.Duration = execResult.Duration
|
||||
result.ExitCode = execResult.ExitCode
|
||||
result.StartTime = execResult.StartTime
|
||||
result.EndTime = execResult.EndTime
|
||||
} else {
|
||||
result.Success = false
|
||||
result.Status = constant.TaskStatusFailed
|
||||
result.StartTime = start
|
||||
result.EndTime = time.Now()
|
||||
result.Duration = result.EndTime.Sub(result.StartTime).Milliseconds()
|
||||
result.Output = rawStr
|
||||
}
|
||||
|
||||
if execErr != nil {
|
||||
result.Error = execErr.Error()
|
||||
if ctx.Err() == context.Canceled {
|
||||
result.Status = constant.TaskStatusCancelled
|
||||
} else if ctx.Err() == context.DeadlineExceeded {
|
||||
result.Status = constant.TaskStatusTimeout
|
||||
}
|
||||
}
|
||||
|
||||
// 6. 执行后事件
|
||||
if s.handler != nil {
|
||||
if execResult != nil {
|
||||
// 只要有执行结果(即使执行失败),都认为是任务完成了(包含输出)
|
||||
s.handler.OnTaskCompleted(req, result)
|
||||
} else if execErr != nil {
|
||||
// 只有在完全没有结果的情况下(如无法启动、Panic等),才认为是任务失败
|
||||
s.handler.OnTaskFailed(req, execErr)
|
||||
}
|
||||
}
|
||||
|
||||
if execErr != nil {
|
||||
s.logger.Errorf("[Scheduler] 任务 %s 执行失败: %v", req.TaskID, execErr)
|
||||
} else {
|
||||
s.logger.Infof("[Scheduler] 执行完成: %s (#%s) [%s] (状态: %s, 耗时: %dms)",
|
||||
req.Name, req.TaskID, req.Type, result.Status, result.Duration)
|
||||
}
|
||||
|
||||
return result, execErr
|
||||
}
|
||||
|
||||
// StopTask 停止正在运行的任务(通过 TaskID,可能会停止多个并发副本)
|
||||
func (s *Scheduler) StopTask(taskID string) bool {
|
||||
s.mu.RLock()
|
||||
cancel, exists := s.runningTasks[taskID]
|
||||
s.mu.RUnlock()
|
||||
|
||||
if exists && cancel != nil {
|
||||
cancel()
|
||||
s.logger.Infof("[Scheduler] 已尝试停止任务 %s", taskID)
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// StopLog 停止正在运行的任务(通过 LogID,精确停止单个执行副本)
|
||||
func (s *Scheduler) StopLog(logID string) bool {
|
||||
s.mu.RLock()
|
||||
cancel, exists := s.runningExecs[logID]
|
||||
s.mu.RUnlock()
|
||||
|
||||
if exists && cancel != nil {
|
||||
cancel()
|
||||
s.logger.Infof("[Scheduler] 已尝试停止任务执行 #%s", logID)
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// GetRunningTaskCount 获取正在运行的任务数量
|
||||
func (s *Scheduler) GetRunningTaskCount() int {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return len(s.runningTasks)
|
||||
}
|
||||
|
||||
// GetRunningTasks 获取所有正在运行的任务 ID
|
||||
func (s *Scheduler) GetRunningTasks() []string {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
ids := make([]string, 0, len(s.runningTasks))
|
||||
for id := range s.runningTasks {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
// Reload 重新加载配置
|
||||
func (s *Scheduler) Reload(config SchedulerConfig) {
|
||||
s.logger.Infof("[Scheduler] 正在重载配置...")
|
||||
|
||||
// 停止现有 workers
|
||||
close(s.stopCh)
|
||||
s.wg.Wait()
|
||||
|
||||
// 更新配置
|
||||
s.mu.Lock()
|
||||
s.config = config
|
||||
s.taskQueue = make(chan *ExecutionRequest, config.QueueSize)
|
||||
s.rateLimiter = time.Tick(config.RateInterval)
|
||||
s.stopCh = make(chan struct{})
|
||||
s.mu.Unlock()
|
||||
|
||||
// 重启 workers
|
||||
s.Start()
|
||||
|
||||
s.logger.Infof("[Scheduler] 配置已重载: workers=%d, queue=%d, rate=%v, strict=%t",
|
||||
config.WorkerCount, config.QueueSize, config.RateInterval, config.StrictQueue)
|
||||
}
|
||||
|
||||
// GetQueueSize 获取当前队列大小
|
||||
func (s *Scheduler) GetQueueSize() int {
|
||||
return len(s.taskQueue)
|
||||
}
|
||||
|
||||
// GetConfig 获取配置
|
||||
func (s *Scheduler) GetConfig() SchedulerConfig {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.config
|
||||
}
|
||||
|
||||
// GetWorkerStatuses 获取所有 Worker 的状态
|
||||
func (s *Scheduler) GetWorkerStatuses() []WorkerStatus {
|
||||
s.workerMu.RLock()
|
||||
defer s.workerMu.RUnlock()
|
||||
// 返回副本防止外部修改
|
||||
statuses := make([]WorkerStatus, len(s.workers))
|
||||
now := time.Now().Unix()
|
||||
for i, w := range s.workers {
|
||||
statuses[i] = w
|
||||
// 在服务端计算运行时间,彻底避免客户端与服务端时钟不一致导致的计算偏差
|
||||
if w.Status == "running" && w.StartTime > 0 {
|
||||
duration := now - w.StartTime
|
||||
if duration < 0 {
|
||||
duration = 0
|
||||
}
|
||||
statuses[i].Duration = duration
|
||||
}
|
||||
}
|
||||
return statuses
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package executor
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
"github.com/engigu/taskpool/internal/systime"
|
||||
"github.com/robfig/cron/v3"
|
||||
)
|
||||
|
||||
type SysCronManager struct {
|
||||
cron *cron.Cron
|
||||
}
|
||||
|
||||
var (
|
||||
sysCronInstance *SysCronManager
|
||||
sysCronOnce sync.Once
|
||||
)
|
||||
|
||||
|
||||
// InitSysCron 初始化系统的内部定时器
|
||||
func InitSysCron(){
|
||||
GetSysCron()
|
||||
}
|
||||
|
||||
// GetSysCron 获取内部系统定时器服务单例
|
||||
func GetSysCron() *SysCronManager {
|
||||
sysCronOnce.Do(func() {
|
||||
// 使用秒级精度,指定为东八区
|
||||
c := cron.New(cron.WithSeconds(), cron.WithLocation(systime.CST))
|
||||
c.Start()
|
||||
sysCronInstance = &SysCronManager{
|
||||
cron: c,
|
||||
}
|
||||
logger.Infof("[SysCron] 内部系统定时管理器已启动")
|
||||
})
|
||||
return sysCronInstance
|
||||
}
|
||||
|
||||
// AddJob 添加内部系统任务,spec为cron表达式(支持 @every 30s 这种快捷方式)
|
||||
func (s *SysCronManager) AddJob(spec string, cmd func()) (cron.EntryID, error) {
|
||||
id, err := s.cron.AddFunc(spec, cmd)
|
||||
if err != nil {
|
||||
logger.Errorf("[SysCron] 无法添加系统任务: %s, err: %v", spec, err)
|
||||
return 0, err
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// AddJobWithRun 立即开启一个协程异步执行一次任务,随后将其加入到系统定时任务中
|
||||
func (s *SysCronManager) AddJobWithRun(spec string, cmd func()) (cron.EntryID, error) {
|
||||
// 立即异步执行一次
|
||||
go func() {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logger.Errorf("[SysCron] 立即执行任务时发生 panic: %v", r)
|
||||
}
|
||||
}()
|
||||
cmd()
|
||||
}()
|
||||
|
||||
// 然后加入定时器
|
||||
return s.AddJob(spec, cmd)
|
||||
}
|
||||
|
||||
// RemoveJob 动态移除指定的系统定时任务
|
||||
func (s *SysCronManager) RemoveJob(id cron.EntryID) {
|
||||
s.cron.Remove(id)
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
package logger
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/systime"
|
||||
"go.uber.org/zap"
|
||||
"go.uber.org/zap/zapcore"
|
||||
)
|
||||
|
||||
var Log *zap.Logger
|
||||
var Sugar *zap.SugaredLogger
|
||||
var atomicLevel zap.AtomicLevel
|
||||
|
||||
// ANSI 颜色代码
|
||||
const (
|
||||
colorReset = "\033[0m"
|
||||
colorRed = "\033[31m" // error, fatal, panic
|
||||
colorYellow = "\033[33m" // warn
|
||||
colorBlue = "\033[36m" // info
|
||||
colorGray = "\033[37m" // debug
|
||||
)
|
||||
|
||||
// customCore 实现 zapcore.Core 以提供与 logrus 一模一样的格式
|
||||
type customCore struct {
|
||||
level zapcore.LevelEnabler
|
||||
writer zapcore.WriteSyncer
|
||||
}
|
||||
|
||||
func (c *customCore) Enabled(l zapcore.Level) bool {
|
||||
return c.level.Enabled(l)
|
||||
}
|
||||
|
||||
func (c *customCore) With(fields []zapcore.Field) zapcore.Core {
|
||||
// 目前忽略字段,以保持与旧 logrus 格式一模一样(旧格式只输出 entry.Message)
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *customCore) Check(ent zapcore.Entry, ce *zapcore.CheckedEntry) *zapcore.CheckedEntry {
|
||||
if c.Enabled(ent.Level) {
|
||||
return ce.AddCore(ent, c)
|
||||
}
|
||||
return ce
|
||||
}
|
||||
|
||||
func (c *customCore) Write(ent zapcore.Entry, fields []zapcore.Field) error {
|
||||
// 统一使用东八区时间
|
||||
timestamp := systime.InCST(ent.Time).Format("2006-01-02 15:04:05")
|
||||
level := strings.ToUpper(ent.Level.String())
|
||||
|
||||
var levelColor string
|
||||
switch ent.Level {
|
||||
case zapcore.DebugLevel:
|
||||
levelColor = colorGray
|
||||
case zapcore.InfoLevel:
|
||||
levelColor = colorBlue
|
||||
case zapcore.WarnLevel:
|
||||
levelColor = colorYellow
|
||||
case zapcore.ErrorLevel, zapcore.DPanicLevel, zapcore.PanicLevel, zapcore.FatalLevel:
|
||||
levelColor = colorRed
|
||||
default:
|
||||
levelColor = colorBlue
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("[%s]%s[%s]%s %s\n", timestamp, levelColor, level, colorReset, ent.Message)
|
||||
_, err := c.writer.Write([]byte(msg))
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *customCore) Sync() error {
|
||||
return c.writer.Sync()
|
||||
}
|
||||
|
||||
func newLogger(output zapcore.WriteSyncer) *zap.Logger {
|
||||
core := &customCore{
|
||||
level: atomicLevel,
|
||||
writer: output,
|
||||
}
|
||||
return zap.New(core)
|
||||
}
|
||||
|
||||
func init() {
|
||||
// 强制设置全局时区为东八区
|
||||
time.Local = systime.CST
|
||||
atomicLevel = zap.NewAtomicLevelAt(zap.InfoLevel)
|
||||
Log = newLogger(zapcore.AddSync(os.Stdout))
|
||||
Sugar = Log.Sugar()
|
||||
}
|
||||
|
||||
// SetupFileOutput 设置文件输出
|
||||
func SetupFileOutput(logDir string) error {
|
||||
if err := os.MkdirAll(logDir, 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
logFile := filepath.Join(logDir, systime.FormatDate(time.Now())+".log")
|
||||
file, err := os.OpenFile(logFile, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0666)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
Log = newLogger(zapcore.AddSync(file))
|
||||
Sugar = Log.Sugar()
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetOutput 直接设置 Log 实例
|
||||
func SetOutput(l *zap.Logger) {
|
||||
Log = l
|
||||
Sugar = l.Sugar()
|
||||
}
|
||||
|
||||
// SetSugar 直接设置 Sugar 实例
|
||||
func SetSugar(s *zap.SugaredLogger) {
|
||||
Sugar = s
|
||||
}
|
||||
|
||||
// SetLevel 设置日志级别
|
||||
func SetLevel(level string) {
|
||||
switch level {
|
||||
case "debug":
|
||||
atomicLevel.SetLevel(zap.DebugLevel)
|
||||
case "info":
|
||||
atomicLevel.SetLevel(zap.InfoLevel)
|
||||
case "warn":
|
||||
atomicLevel.SetLevel(zap.WarnLevel)
|
||||
case "error":
|
||||
atomicLevel.SetLevel(zap.ErrorLevel)
|
||||
default:
|
||||
atomicLevel.SetLevel(zap.InfoLevel)
|
||||
}
|
||||
}
|
||||
|
||||
// 便捷方法
|
||||
func Debug(args ...interface{}) { Sugar.Debug(args...) }
|
||||
func Info(args ...interface{}) { Sugar.Info(args...) }
|
||||
func Warn(args ...interface{}) { Sugar.Warn(args...) }
|
||||
func Error(args ...interface{}) { Sugar.Error(args...) }
|
||||
func Fatal(args ...interface{}) { Sugar.Fatal(args...) }
|
||||
|
||||
func Debugf(format string, args ...interface{}) { Sugar.Debugf(format, args...) }
|
||||
func Infof(format string, args ...interface{}) { Sugar.Infof(format, args...) }
|
||||
func Warnf(format string, args ...interface{}) { Sugar.Warnf(format, args...) }
|
||||
func Errorf(format string, args ...interface{}) { Sugar.Errorf(format, args...) }
|
||||
func Fatalf(format string, args ...interface{}) { Sugar.Fatalf(format, args...) }
|
||||
|
||||
// WithField 带字段的日志
|
||||
func WithField(key string, value interface{}) *zap.SugaredLogger {
|
||||
return Sugar.With(key, value)
|
||||
}
|
||||
|
||||
// WithFields 带多个字段的日志
|
||||
func WithFields(fields map[string]interface{}) *zap.SugaredLogger {
|
||||
f := make([]interface{}, 0, len(fields)*2)
|
||||
for k, v := range fields {
|
||||
f = append(f, k, v)
|
||||
}
|
||||
return Sugar.With(f...)
|
||||
}
|
||||
|
||||
// SchedulerLogger 兼容 internal/executor 的日志接口
|
||||
type SchedulerLogger struct{}
|
||||
|
||||
func (s *SchedulerLogger) Infof(format string, args ...interface{}) {
|
||||
Sugar.Infof(format, args...)
|
||||
}
|
||||
func (s *SchedulerLogger) Warnf(format string, args ...interface{}) {
|
||||
Sugar.Warnf(format, args...)
|
||||
}
|
||||
func (s *SchedulerLogger) Errorf(format string, args ...interface{}) {
|
||||
Sugar.Errorf(format, args...)
|
||||
}
|
||||
|
||||
// NewSchedulerLogger 创建一个兼容 executor.SchedulerLogger 的实例
|
||||
func NewSchedulerLogger() *SchedulerLogger {
|
||||
return &SchedulerLogger{}
|
||||
}
|
||||
@@ -0,0 +1,329 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"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/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// AuthRequired 认证中间件
|
||||
func AuthRequired() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 基础的 CSRF 防护:校验 Origin/Referer (针对非 GET 请求)
|
||||
if c.Request.Method != http.MethodGet && c.Request.Method != http.MethodOptions && c.Request.Method != http.MethodHead {
|
||||
origin := c.GetHeader("Origin")
|
||||
if origin == "" {
|
||||
origin = c.GetHeader("Referer")
|
||||
}
|
||||
// 如果有 Origin 且不匹配则拒绝(实际部署时应配置允许的 Origin)
|
||||
if origin != "" && !utils.CheckWSOrigin(c.Request) {
|
||||
utils.Forbidden(c, "CSRF 校验失败: 非法的请求来源")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 检查是否携带互联 Token(支持跨面板远程全接口调用)
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if authHeader != "" {
|
||||
tokenStr := strings.TrimSpace(strings.TrimPrefix(authHeader, "Bearer "))
|
||||
if tokenStr != "" {
|
||||
settingsSvc := services.NewSettingsService()
|
||||
interconnectToken := settingsSvc.Get(constant.SectionSite, constant.KeyInterconnectToken)
|
||||
parentToken := settingsSvc.Get(constant.SectionInterconnect, constant.KeyInterconnectParentToken)
|
||||
|
||||
isMatched := false
|
||||
h1 := sha256.Sum256([]byte(tokenStr))
|
||||
|
||||
if interconnectToken != "" {
|
||||
h2 := sha256.Sum256([]byte(interconnectToken))
|
||||
if subtle.ConstantTimeCompare(h1[:], h2[:]) == 1 {
|
||||
isMatched = true
|
||||
}
|
||||
}
|
||||
|
||||
if !isMatched && parentToken != "" {
|
||||
h2 := sha256.Sum256([]byte(parentToken))
|
||||
if subtle.ConstantTimeCompare(h1[:], h2[:]) == 1 {
|
||||
isMatched = true
|
||||
}
|
||||
}
|
||||
|
||||
if isMatched {
|
||||
// 模拟 Admin 角色
|
||||
var adminUser models.User
|
||||
res := database.DB.Where("role = ?", constant.AdminRole).Limit(1).Find(&adminUser)
|
||||
if res.Error == nil && res.RowsAffected > 0 {
|
||||
c.Set("userID", adminUser.ID)
|
||||
c.Set("username", adminUser.Username)
|
||||
c.Set("role", adminUser.Role)
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
token, err := c.Cookie(constant.CookieName)
|
||||
if err != nil || token == "" {
|
||||
utils.Unauthorized(c, "请先登录")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 验证 token
|
||||
userID, username, tokenVersion, err := utils.ParseToken(token, constant.Secret)
|
||||
if err != nil {
|
||||
utils.Unauthorized(c, "登录已过期,请重新登录")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 安全增强:校验数据库中该用户的 ID 是否与 Token 一致,并验证 TokenVersion
|
||||
var user models.User
|
||||
res := database.DB.Where("username = ?", username).Limit(1).Find(&user)
|
||||
if res.Error != nil || res.RowsAffected == 0 || user.ID != userID || user.TokenVersion != tokenVersion {
|
||||
utils.Unauthorized(c, "会话失效,请重新登录")
|
||||
ClearAuthCookie(c)
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 将用户信息存入上下文 (必须使用数据库中的最新 ID)
|
||||
c.Set("userID", user.ID)
|
||||
c.Set("username", user.Username)
|
||||
c.Set("role", user.Role)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// AdminRequired 管理员权限认证中间件
|
||||
func AdminRequired() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
role, exists := c.Get("role")
|
||||
if !exists || role != constant.AdminRole {
|
||||
utils.Forbidden(c, "需要管理员权限")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// OpenapiRequired OpenAPI 认证中间件
|
||||
func OpenapiRequired() gin.HandlerFunc {
|
||||
settingsSvc := services.NewSettingsService()
|
||||
return func(c *gin.Context) {
|
||||
if checkOpenapiToken(c, settingsSvc) {
|
||||
return
|
||||
}
|
||||
utils.Unauthorized(c, "无效的 OpenAPI 令牌")
|
||||
c.Abort()
|
||||
}
|
||||
}
|
||||
|
||||
// checkOpenapiToken 校验 OpenAPI Token
|
||||
// 返回 true 表示校验通过并已放行请求
|
||||
func checkOpenapiToken(c *gin.Context, settingsSvc *services.SettingsService) bool {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if authHeader == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
// 提取 token:支持 "Bearer <token>" 和直接 "<token>" 两种格式
|
||||
var openapiToken string
|
||||
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
|
||||
// 标准格式:Bearer <token>
|
||||
openapiToken = authHeader[7:]
|
||||
} else {
|
||||
// 直接使用 token
|
||||
openapiToken = authHeader
|
||||
}
|
||||
|
||||
// Token 不能为空
|
||||
if openapiToken == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
siteConfig := settingsSvc.GetSection(constant.SectionSite)
|
||||
tokenJson, ok := siteConfig[constant.KeyOpenapiToken]
|
||||
if !ok || tokenJson == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
var tokenConfig vo.TokenConfig
|
||||
if err := json.Unmarshal([]byte(tokenJson), &tokenConfig); err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
// 校验开启状态
|
||||
if !tokenConfig.Enabled {
|
||||
return false
|
||||
}
|
||||
|
||||
if tokenConfig.Token == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
// 使用恒定时间比较防止时序攻击
|
||||
h1 := sha256.Sum256([]byte(openapiToken))
|
||||
h2 := sha256.Sum256([]byte(tokenConfig.Token))
|
||||
if subtle.ConstantTimeCompare(h1[:], h2[:]) != 1 {
|
||||
return false
|
||||
}
|
||||
|
||||
// 检查过期时间
|
||||
if tokenConfig.ExpireAt != "" {
|
||||
expireDate, err := time.Parse("2006-01-02", tokenConfig.ExpireAt)
|
||||
if err == nil {
|
||||
expireDate = expireDate.Add(23*time.Hour + 59*time.Minute + 59*time.Second)
|
||||
if time.Now().After(expireDate) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 模拟 Admin 角色
|
||||
var adminUser models.User
|
||||
res := database.DB.Where("role = ?", "admin").Limit(1).Find(&adminUser)
|
||||
if res.Error != nil || res.RowsAffected == 0 {
|
||||
utils.Unauthorized(c, "未找到管理员账户,OpenAPI Token 校验失败")
|
||||
c.Abort()
|
||||
return true
|
||||
}
|
||||
|
||||
c.Set("userID", adminUser.ID)
|
||||
c.Set("username", adminUser.Username)
|
||||
c.Set("role", adminUser.Role)
|
||||
c.Next()
|
||||
return true
|
||||
}
|
||||
|
||||
// SetAuthCookie 设置认证 Cookie,expireDays 为过期天数
|
||||
func SetAuthCookie(c *gin.Context, token string, expireDays int) {
|
||||
maxAge := 86400 * expireDays
|
||||
// 增加 SameSite=Lax 和 Secure 属性(如果环境支持,这里暂时设为 false,但生产建议 true)
|
||||
c.SetSameSite(http.SameSiteLaxMode)
|
||||
c.SetCookie(constant.CookieName, token, maxAge, "/", "", false, true)
|
||||
}
|
||||
|
||||
// ClearAuthCookie 清除认证 Cookie
|
||||
func ClearAuthCookie(c *gin.Context) {
|
||||
c.SetCookie(constant.CookieName, "", -1, "/", "", false, true)
|
||||
}
|
||||
|
||||
// SwaggerAuth Swagger 认证中间件 (Basic Auth)
|
||||
func SwaggerAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
settingsSvc := services.NewSettingsService()
|
||||
siteConfig := settingsSvc.GetSection(constant.SectionSite)
|
||||
tokenJson := siteConfig[constant.KeyOpenapiToken]
|
||||
|
||||
if tokenJson == "" {
|
||||
c.Status(http.StatusNotFound)
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
var tokenConfig vo.TokenConfig
|
||||
if err := json.Unmarshal([]byte(tokenJson), &tokenConfig); err != nil {
|
||||
c.Status(http.StatusNotFound)
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 必须开启鉴权开关
|
||||
if !tokenConfig.Enabled {
|
||||
c.Status(http.StatusNotFound)
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 检查过期时间
|
||||
if tokenConfig.ExpireAt != "" {
|
||||
expire, err := time.ParseInLocation("2006/01/02", tokenConfig.ExpireAt, time.Local)
|
||||
if err == nil {
|
||||
// 包含当天,所以设置到当天 23:59:59
|
||||
expire = expire.Add(24*time.Hour - time.Second)
|
||||
if time.Now().After(expire) {
|
||||
c.Status(http.StatusNotFound)
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 获取请求中携带的凭证
|
||||
// 1. URL 参数 token
|
||||
// 2. Cookie 中的 openapi_token
|
||||
// 3. HTTP Basic Auth
|
||||
tokenQuery := c.Query("token")
|
||||
tokenCookie, _ := c.Cookie("openapi_token")
|
||||
_, password, hasAuth := c.Request.BasicAuth()
|
||||
|
||||
var providedToken string
|
||||
if tokenQuery != "" {
|
||||
providedToken = tokenQuery
|
||||
} else if tokenCookie != "" {
|
||||
providedToken = tokenCookie
|
||||
} else if hasAuth {
|
||||
providedToken = password
|
||||
}
|
||||
|
||||
// 检查提供的 token 是否匹配
|
||||
if providedToken != "" {
|
||||
h1 := sha256.Sum256([]byte(providedToken))
|
||||
h2 := sha256.Sum256([]byte(tokenConfig.Token))
|
||||
if subtle.ConstantTimeCompare(h1[:], h2[:]) == 1 {
|
||||
// 如果是通过 url 参数进来的,自动将其种入 Cookie,便于后续加载静态资源 (如 json)
|
||||
if tokenQuery != "" {
|
||||
c.SetCookie("openapi_token", providedToken, 86400, "/openapi", "", false, false)
|
||||
}
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 验证失败,不再返回 WWW-Authenticate 头触发浏览器反人类原生弹窗
|
||||
// 我们返回 401 的 JSON 或纯文本结构,以便由调用方自行接管鉴权逻辑
|
||||
c.JSON(http.StatusUnauthorized, gin.H{
|
||||
"code": 401,
|
||||
"msg": "OpenAPI 访问未授权或 Token 错误",
|
||||
})
|
||||
c.Abort()
|
||||
}
|
||||
}
|
||||
|
||||
// LocalhostOnly 仅允许本地回环地址访问,并进行简单的内部凭证校验
|
||||
func LocalhostOnly() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
ip := c.ClientIP()
|
||||
if ip != "127.0.0.1" && ip != "::1" {
|
||||
utils.BadRequest(c, "仅允许本地访问")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 简单的内部通信认证
|
||||
token := c.GetHeader("X-Internal-Token")
|
||||
if token == "" || token != constant.Secret {
|
||||
utils.Unauthorized(c, "无效的内部调用凭证")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// GinLogger 返回使用 logrus 的 Gin 日志中间件
|
||||
func GinLogger() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
start := time.Now()
|
||||
path := c.Request.URL.Path
|
||||
query := c.Request.URL.RawQuery
|
||||
|
||||
c.Next()
|
||||
|
||||
latency := time.Since(start)
|
||||
status := c.Writer.Status()
|
||||
clientIP := c.ClientIP()
|
||||
method := c.Request.Method
|
||||
|
||||
if query != "" {
|
||||
path = path + "?" + query
|
||||
}
|
||||
|
||||
// 格式化延迟时间
|
||||
var latencyStr string
|
||||
if latency < time.Millisecond {
|
||||
latencyStr = fmt.Sprintf("%dµs", latency.Microseconds())
|
||||
} else if latency < time.Second {
|
||||
latencyStr = fmt.Sprintf("%dms", latency.Milliseconds())
|
||||
} else {
|
||||
latencyStr = fmt.Sprintf("%.2fs", latency.Seconds())
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("[HTTP] %d %s %s %s [%s]", status, method, path, latencyStr, clientIP)
|
||||
|
||||
if status >= 500 {
|
||||
logger.Error(msg)
|
||||
} else if status >= 400 {
|
||||
logger.Warn(msg)
|
||||
} else {
|
||||
logger.Info(msg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GinRecovery 返回使用 logrus 的 Gin 恢复中间件
|
||||
func GinRecovery() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
logger.Errorf("[HTTP] Panic: %v | %s", err, c.Request.URL.Path)
|
||||
c.AbortWithStatus(500)
|
||||
}
|
||||
}()
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// NotifyTokenAuth 通知 Token 认证中间件
|
||||
func NotifyTokenAuth() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token := c.GetHeader("notify-token")
|
||||
if token == "" {
|
||||
utils.Unauthorized(c, "缺少通知 Token")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 从 settings 表读取配置的通知 Token
|
||||
settingsService := services.NewSettingsService()
|
||||
savedToken := settingsService.Get(constant.SectionNotify, constant.KeyNotifyToken)
|
||||
|
||||
if savedToken == "" {
|
||||
utils.Unauthorized(c, "通知 Token 未配置")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
if !strings.EqualFold(token, savedToken) {
|
||||
utils.Unauthorized(c, "通知 Token 无效")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"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"
|
||||
"github.com/engigu/taskpool/internal/tunnel"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
var travelProxyClient = &http.Client{
|
||||
Timeout: 10 * time.Second,
|
||||
}
|
||||
|
||||
func TravelProxyMiddleware() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 1. 检查是否存在 active_interconnect_node_id Cookie
|
||||
nodeID, err := c.Cookie(constant.CookieActiveInterconnectNodeID)
|
||||
if err != nil || nodeID == "" {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
// 2. 白名单放行:反向隧道建立连接端点必须直达本机,不能二次代理
|
||||
path := c.Request.URL.Path
|
||||
if strings.HasPrefix(path, "/api/v1/interconnect/tunnel") {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
// 3. 查询数据库中节点信息
|
||||
var node models.InterconnectNode
|
||||
if err := database.DB.Where("id = ?", nodeID).First(&node).Error; err != nil {
|
||||
// 节点不存在,说明 Cookie 无效,清除并放行
|
||||
c.SetCookie(constant.CookieActiveInterconnectNodeID, "", -1, "/", "", false, false)
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
// 4. 准备代理的路径(若主节点配置了 URLPrefix,需要剥离)
|
||||
cfg := services.GetConfig()
|
||||
urlPrefix := strings.TrimSuffix(cfg.Server.URLPrefix, "/")
|
||||
targetPath := path
|
||||
if urlPrefix != "" {
|
||||
targetPath = strings.TrimPrefix(targetPath, urlPrefix)
|
||||
}
|
||||
if !strings.HasPrefix(targetPath, "/") {
|
||||
targetPath = "/" + targetPath
|
||||
}
|
||||
|
||||
// 5. 执行代理转发
|
||||
if strings.HasPrefix(node.URL, "tunnel://") {
|
||||
// 走 WebSocket 逆向 Yamux 隧道
|
||||
err := tunnel.ProxyHTTP(nodeID, c, targetPath)
|
||||
if err != nil {
|
||||
// 如果是网页 HTML 导航请求,提供友好降级返回主节点
|
||||
if strings.Contains(c.GetHeader("Accept"), "text/html") {
|
||||
c.SetCookie(constant.CookieActiveInterconnectNodeID, "", -1, "/", "", false, false)
|
||||
c.Header("Content-Type", "text/html; charset=utf-8")
|
||||
c.String(200, `<p>与子节点连接失败,正在自动返回主节点...</p><script>document.cookie="` + constant.CookieActiveInterconnectNodeID + `=; expires=Thu, 01 Jan 1970 00:00:00 UTC; path=/;"; window.location.href="/";</script>`)
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
c.JSON(502, gin.H{"code": 502, "msg": "与子节点逆向隧道通信异常: " + err.Error()})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 走普通 HTTP 直连代理
|
||||
targetURL := strings.TrimRight(node.URL, "/") + targetPath
|
||||
if c.Request.URL.RawQuery != "" {
|
||||
targetURL += "?" + c.Request.URL.RawQuery
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(c.Request.Method, targetURL, c.Request.Body)
|
||||
if err != nil {
|
||||
c.JSON(500, gin.H{"code": 500, "msg": "Failed to create proxy request"})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 复制请求头
|
||||
req.Header = c.Request.Header.Clone()
|
||||
|
||||
// 覆盖认证授权 Header 确保子节点鉴权通过
|
||||
if node.Token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+node.Token)
|
||||
}
|
||||
// 移除 Cookie 头部防干扰
|
||||
req.Header.Del("Cookie")
|
||||
req.Header.Set("X-Tunnel-Proxy", "true")
|
||||
|
||||
resp, err := travelProxyClient.Do(req)
|
||||
if err != nil {
|
||||
// 如果是客户端自己主动取消了请求(例如连续刷新、关闭网页等),直接退出,不应视为子节点离线而执行退回主节点的操作
|
||||
if c.Request.Context().Err() == context.Canceled || strings.Contains(err.Error(), "context canceled") {
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
if strings.Contains(c.GetHeader("Accept"), "text/html") {
|
||||
c.SetCookie(constant.CookieActiveInterconnectNodeID, "", -1, "/", "", false, false)
|
||||
c.Header("Content-Type", "text/html; charset=utf-8")
|
||||
c.String(200, `<p>与子节点连接失败,正在自动返回主节点...</p><script>document.cookie="` + constant.CookieActiveInterconnectNodeID + `=; expires=Thu, 01 Jan 1970 00:00:00 UTC; path=/;"; window.location.href="/";</script>`)
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
c.JSON(502, gin.H{"code": 502, "msg": "无法连接至目标子节点: " + err.Error()})
|
||||
c.Abort()
|
||||
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)
|
||||
c.Abort()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// AgentSchedulerConfig Agent 调度器配置
|
||||
type AgentSchedulerConfig struct {
|
||||
WorkerCount int `json:"worker_count"`
|
||||
QueueSize int `json:"queue_size"`
|
||||
RateInterval time.Duration `json:"rate_interval"`
|
||||
Verbose bool `json:"verbose"`
|
||||
StrictQueue bool `json:"strict_queue"`
|
||||
}
|
||||
|
||||
// Value 序列化为数据库字符串
|
||||
func (c AgentSchedulerConfig) Value() (driver.Value, error) {
|
||||
bytes, err := json.Marshal(c)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return string(bytes), nil
|
||||
}
|
||||
|
||||
// Scan 反序列化数据库字符串为结构体
|
||||
func (c *AgentSchedulerConfig) Scan(value interface{}) error {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
bytes, ok := value.([]byte)
|
||||
if !ok {
|
||||
str, ok := value.(string)
|
||||
if !ok {
|
||||
return errors.New("invalid type for AgentSchedulerConfig")
|
||||
}
|
||||
bytes = []byte(str)
|
||||
}
|
||||
return json.Unmarshal(bytes, c)
|
||||
}
|
||||
|
||||
// Agent 远程执行代理
|
||||
type Agent struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Name string `json:"name" gorm:"size:100;not null"` // Agent 名称
|
||||
Token string `json:"token" gorm:"size:64;index"` // 认证 Token(可重复使用)
|
||||
MachineID string `json:"machine_id" gorm:"size:64;uniqueIndex"` // 机器识别码(唯一)
|
||||
Description string `json:"description" gorm:"size:255"` // 描述
|
||||
Status string `json:"status" gorm:"size:20;default:'pending';index"` // 状态: constant.AgentStatusOnline, constant.AgentStatusOffline
|
||||
LastSeen *LocalTime `json:"last_seen"` // 最后心跳时间
|
||||
IP string `json:"ip" gorm:"size:45"` // Agent IP 地址
|
||||
Version string `json:"version" gorm:"size:50"` // Agent 版本
|
||||
BuildTime string `json:"build_time" gorm:"size:30"` // Agent 构建时间
|
||||
Hostname string `json:"hostname" gorm:"size:100"` // Agent 主机名
|
||||
OS string `json:"os" gorm:"size:20"` // 操作系统
|
||||
Arch string `json:"arch" gorm:"size:20"` // 架构
|
||||
ForceUpdate bool `json:"force_update" gorm:"default:false"` // 强制更新标志
|
||||
Enabled *bool `json:"enabled" gorm:"default:true"` // 是否启用
|
||||
SchedulerConfig AgentSchedulerConfig `json:"scheduler_config" gorm:"type:text"` // 调度配置,以 JSON 字符串形式存储在 Text 类型字段中
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (Agent) TableName() string {
|
||||
return constant.TablePrefix + "agents"
|
||||
}
|
||||
|
||||
// AgentToken Agent 令牌
|
||||
type AgentToken struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Token string `json:"token" gorm:"size:64;uniqueIndex;not null"` // 令牌
|
||||
Remark string `json:"remark" gorm:"size:255"` // 备注
|
||||
MaxUses int `json:"max_uses" gorm:"default:0"` // 最大使用次数,0 表示无限制
|
||||
UsedCount int `json:"used_count" gorm:"default:0"` // 已使用次数
|
||||
ExpiresAt *LocalTime `json:"expires_at"` // 过期时间,null 表示永不过期
|
||||
Enabled *bool `json:"enabled" gorm:"default:true"` // 是否启用
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (AgentToken) TableName() string {
|
||||
return constant.TablePrefix + "tokens"
|
||||
}
|
||||
|
||||
// AgentTask Agent 任务配置(用于下发给 Agent)
|
||||
type AgentTask struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Command string `json:"command"`
|
||||
PreCommand string `json:"pre_command"`
|
||||
PostCommand string `json:"post_command"`
|
||||
Schedule string `json:"schedule"`
|
||||
Timeout int `json:"timeout"`
|
||||
WorkDir string `json:"work_dir"`
|
||||
Envs string `json:"envs"`
|
||||
Languages []map[string]string `json:"languages"`
|
||||
RandomRange int `json:"random_range"`
|
||||
Secrets []string `json:"secrets"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
func (t AgentTask) GetID() string {
|
||||
return t.ID
|
||||
}
|
||||
|
||||
func (t AgentTask) GetName() string {
|
||||
return t.Name
|
||||
}
|
||||
|
||||
func (t AgentTask) GetCommand() string {
|
||||
return t.Command
|
||||
}
|
||||
|
||||
func (t AgentTask) GetPreCommand() string {
|
||||
return t.PreCommand
|
||||
}
|
||||
|
||||
func (t AgentTask) GetPostCommand() string {
|
||||
return t.PostCommand
|
||||
}
|
||||
|
||||
func (t AgentTask) GetSchedule() string {
|
||||
return t.Schedule
|
||||
}
|
||||
|
||||
func (t AgentTask) GetRandomRange() int {
|
||||
return t.RandomRange
|
||||
}
|
||||
|
||||
func (t AgentTask) GetSecrets() []string {
|
||||
return t.Secrets
|
||||
}
|
||||
|
||||
// AgentTaskResult Agent 上报的任务执行结果
|
||||
type AgentTaskResult struct {
|
||||
TaskID string `json:"task_id"`
|
||||
LogID string `json:"log_id"`
|
||||
AgentID string `json:"agent_id"`
|
||||
Command string `json:"command"`
|
||||
Output string `json:"output"`
|
||||
Error string `json:"error"` // 额外的系统错误信息
|
||||
Status string `json:"status"` // success, failed
|
||||
Duration int64 `json:"duration"` // 耗时(毫秒)
|
||||
ExitCode int `json:"exit_code"`
|
||||
StartTime int64 `json:"start_time"` // Unix 时间戳
|
||||
EndTime int64 `json:"end_time"` // Unix 时间戳
|
||||
}
|
||||
|
||||
// AgentRegisterRequest Agent 注册请求
|
||||
type AgentRegisterRequest struct {
|
||||
Name string `json:"name"`
|
||||
Hostname string `json:"hostname"`
|
||||
Version string `json:"version"`
|
||||
BuildTime string `json:"build_time"`
|
||||
Token string `json:"token"` // 注册令牌
|
||||
MachineID string `json:"machine_id"` // 机器识别码
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// AppLog 统一应用日志与通知记录
|
||||
type AppLog struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Category string `json:"category" gorm:"size:50;index;not null"` // 大类:constant.LogCategorySystemNotice(系统通知), constant.LogCategoryPushLog(推送记录) 等
|
||||
Title string `json:"title" gorm:"size:255"` // 消息标题
|
||||
Content BigText `json:"content"` // 详细内容/Payload
|
||||
Level string `json:"level" gorm:"size:20;index"` // 级别:constant.LogLevelInfo, constant.LogLevelWarning, constant.LogLevelError
|
||||
Status string `json:"status" gorm:"size:20;index"` // 状态:系统通知为 constant.LogStatusRead/constant.LogStatusUnread,推送为 constant.LogStatusSuccess/constant.LogStatusFailed
|
||||
RefID string `json:"ref_id" gorm:"size:50;index"` // 关联对象ID(选填,比如绑定的通知渠道ID、任务ID等)
|
||||
ErrorMsg BigText `json:"error_msg"` // 执行错误信息详情
|
||||
CreatedAt LocalTime `json:"created_at" gorm:"index"`
|
||||
ReadAt *LocalTime `json:"read_at"` // 已读时间(仅对通知生效)
|
||||
ChannelName string `json:"channel_name" gorm:"-"` // 推送记录的关联渠道名称(动态查询)
|
||||
}
|
||||
|
||||
func (AppLog) TableName() string {
|
||||
return constant.TablePrefix + "app_logs"
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/schema"
|
||||
)
|
||||
|
||||
// BigText 自定义大数据文本类型,自动处理跨数据库类型差异
|
||||
// MySQL: LONGTEXT (4GB)
|
||||
// PostgreSQL: TEXT
|
||||
// SQLite: TEXT
|
||||
type BigText string
|
||||
|
||||
func (BigText) GormDBDataType(db *gorm.DB, field *schema.Field) string {
|
||||
switch db.Dialector.Name() {
|
||||
case "mysql":
|
||||
return "LONGTEXT"
|
||||
case "postgres":
|
||||
return "TEXT"
|
||||
case "sqlite":
|
||||
return "TEXT"
|
||||
}
|
||||
return "TEXT"
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// DataRelation 通用数据关联表
|
||||
type DataRelation struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
DataID string `json:"data_id" gorm:"size:20;index;not null"`
|
||||
RelateID string `json:"relate_id" gorm:"size:20;index;not null"`
|
||||
Type string `json:"type" gorm:"size:50;index;not null"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (DataRelation) TableName() string {
|
||||
return constant.TablePrefix + "data_relations"
|
||||
}
|
||||
|
||||
// DataStorage 通用数据存储表
|
||||
type DataStorage struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Type string `json:"type" gorm:"size:50;index;not null"`
|
||||
Name string `json:"name" gorm:"size:255;index;not null"`
|
||||
Data BigText `json:"data"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (DataStorage) TableName() string {
|
||||
return constant.TablePrefix + "data_storages"
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// Dependency 依赖包模型
|
||||
type Dependency struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Name string `json:"name" gorm:"size:100;not null"`
|
||||
Version string `json:"version" gorm:"size:50"`
|
||||
Language string `json:"language" gorm:"size:100;index"` // 关联语言 (node, python...)
|
||||
LangVersion string `json:"lang_version" gorm:"size:100;index"` // 关联语言版本
|
||||
Remark string `json:"remark" gorm:"size:255"`
|
||||
Log BigText `json:"log"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (Dependency) TableName() string {
|
||||
return constant.TablePrefix + "deps"
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// EnvironmentVariable represents an environment variable
|
||||
type EnvironmentVariable struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Name string `json:"name" gorm:"size:255;not null"`
|
||||
Value BigText `json:"value"`
|
||||
Remark string `json:"remark" gorm:"size:500"`
|
||||
Type string `json:"type" gorm:"size:20;default:'normal'"`
|
||||
Hidden *bool `json:"hidden" gorm:"default:true"`
|
||||
Enabled *bool `json:"enabled" gorm:"default:true"`
|
||||
UserID string `json:"user_id" gorm:"size:20;index"`
|
||||
Tags string `json:"-" gorm:"-"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (EnvironmentVariable) TableName() string {
|
||||
return constant.TablePrefix + "envs"
|
||||
}
|
||||
|
||||
// Script represents a script file
|
||||
type Script struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Name string `json:"name" gorm:"size:255;not null"`
|
||||
Content BigText `json:"content"`
|
||||
UserID string `json:"user_id" gorm:"size:20;index"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (Script) TableName() string {
|
||||
return constant.TablePrefix + "scripts"
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
package models
|
||||
|
||||
import "time"
|
||||
|
||||
// ExportData 全量或部分业务数据的导出/导入结构
|
||||
type ExportData struct {
|
||||
Version string `json:"version"`
|
||||
ExportAt LocalTime `json:"export_at"`
|
||||
Tasks []Task `json:"tasks"`
|
||||
Envs []EnvironmentVariable `json:"envs"`
|
||||
Tags []DataStorage `json:"tags"`
|
||||
Bindings []NotifyBinding `json:"bindings"`
|
||||
}
|
||||
|
||||
// NewExportData 创建一个导出数据对象
|
||||
func NewExportData() *ExportData {
|
||||
return &ExportData{
|
||||
Version: "1.0",
|
||||
ExportAt: LocalTime(time.Now()),
|
||||
Tasks: make([]Task, 0),
|
||||
Envs: make([]EnvironmentVariable, 0),
|
||||
Tags: make([]DataStorage, 0),
|
||||
Bindings: make([]NotifyBinding, 0),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// NodeMetrics 表示互联节点的性能指标,使用 JSON 存储
|
||||
type NodeMetrics struct {
|
||||
CPUPercent float64 `json:"cpu_percent"`
|
||||
MemPercent float64 `json:"mem_percent"`
|
||||
DiskPercent float64 `json:"disk_percent"`
|
||||
TxBytes uint64 `json:"tx_bytes,omitempty"`
|
||||
RxBytes uint64 `json:"rx_bytes,omitempty"`
|
||||
}
|
||||
|
||||
// InterconnectNode represents a connected remote panel
|
||||
type InterconnectNode struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Name string `json:"name" gorm:"size:255;not null"`
|
||||
URL string `json:"url" gorm:"size:255;not null"`
|
||||
Token string `json:"token" gorm:"size:255"`
|
||||
Remark string `json:"remark" gorm:"size:500"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
Status string `json:"status" gorm:"size:50"` // online / offline
|
||||
Metrics NodeMetrics `json:"metrics" gorm:"serializer:json"`
|
||||
LastHeartbeatAt *LocalTime `json:"last_heartbeat_at"`
|
||||
}
|
||||
|
||||
func (InterconnectNode) TableName() string {
|
||||
return constant.TablePrefix + "interconnect_nodes"
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
type Language struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Plugin string `json:"plugin" gorm:"size:100;not null;index"`
|
||||
Version string `json:"version" gorm:"size:100;not null;index"`
|
||||
InstallPath string `json:"install_path" gorm:"size:255"`
|
||||
Source string `json:"source" gorm:"size:255"`
|
||||
InstalledAt *LocalTime `json:"installed_at"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (Language) TableName() string {
|
||||
return constant.TablePrefix + "languages"
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// NotifyBinding 事件绑定表
|
||||
type NotifyBinding struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Type string `json:"type" gorm:"size:20;not null;index"` // system 或 task
|
||||
Event string `json:"event" gorm:"size:50;not null;index"` // 事件类型
|
||||
WayID string `json:"way_id" gorm:"size:20;not null;index"` // 通知渠道ID
|
||||
DataID string `json:"data_id" gorm:"size:20;index"` // 关联ID,系统事件为空,任务事件为任务ID
|
||||
Extra BigText `json:"extra"` // 额外配置(如是否开启日志推送等,对应 BindingExtra 结构)
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
// BindingExtra 存储在 Extra 字段中的 JSON 配置
|
||||
type BindingExtra struct {
|
||||
EnableLog bool `json:"enable_log"`
|
||||
LogLimit int `json:"log_limit"` // 日志字数限制,默认 1000
|
||||
}
|
||||
|
||||
func (NotifyBinding) TableName() string {
|
||||
return constant.TablePrefix + "notify_bindings"
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// NotifyWay 消息推送渠道
|
||||
type NotifyWay struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Name string `json:"name" gorm:"size:100;not null"`
|
||||
Type string `json:"type" gorm:"size:50;not null;index"`
|
||||
Config BigText `json:"config"`
|
||||
Enabled *bool `json:"enabled" gorm:"default:true;index"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (NotifyWay) TableName() string {
|
||||
return constant.TablePrefix + "notify_ways"
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// SendStats 任务执行统计
|
||||
type SendStats struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
TaskID string `json:"task_id" gorm:"size:20;uniqueIndex:idx_task_day_status"`
|
||||
Day string `json:"day" gorm:"size:10;uniqueIndex:idx_task_day_status"` // 格式: 2006-01-02
|
||||
Status string `json:"status" gorm:"size:20;uniqueIndex:idx_task_day_status"`
|
||||
Num int `json:"num" gorm:"default:0"`
|
||||
}
|
||||
|
||||
func (SendStats) TableName() string {
|
||||
return constant.TablePrefix + "send_stats"
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// Setting 系统设置
|
||||
type Setting struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Section string `json:"section" gorm:"size:50;not null;index:idx_section_key"`
|
||||
Key string `json:"key" gorm:"size:100;not null;index:idx_section_key"`
|
||||
Value BigText `json:"value"`
|
||||
}
|
||||
|
||||
func (Setting) TableName() string {
|
||||
return constant.TablePrefix + "settings"
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// TaskLanguages 自定义语言配置列表类型,处理 JSON 序列化
|
||||
type TaskLanguages []map[string]string
|
||||
|
||||
func (t TaskLanguages) Value() (driver.Value, error) {
|
||||
if t == nil {
|
||||
return "[]", nil
|
||||
}
|
||||
b, err := json.Marshal(t)
|
||||
return string(b), err
|
||||
}
|
||||
|
||||
func (t *TaskLanguages) Scan(v interface{}) error {
|
||||
if v == nil {
|
||||
*t = nil
|
||||
return nil
|
||||
}
|
||||
var data []byte
|
||||
switch s := v.(type) {
|
||||
case string:
|
||||
data = []byte(s)
|
||||
case []byte:
|
||||
data = s
|
||||
default:
|
||||
return fmt.Errorf("invalid type for TaskLanguages: %T", v)
|
||||
}
|
||||
return json.Unmarshal(data, t)
|
||||
}
|
||||
|
||||
// CleanConfig 清理配置结构
|
||||
type CleanConfig struct {
|
||||
Type string `json:"type"` // "day" 或 "count"
|
||||
Keep int `json:"keep"` // 保留天数或条数
|
||||
}
|
||||
|
||||
// RepoConfig 仓库同步配置
|
||||
type RepoConfig struct {
|
||||
SourceType string `json:"source_type"` // url 或 git
|
||||
SourceURL string `json:"source_url"` // 源地址
|
||||
TargetPath string `json:"target_path"` // 目标路径
|
||||
Branch string `json:"branch"` // Git 分支
|
||||
SparsePath string `json:"sparse_path"` // 稀疏检出路径(仅拉取指定目录或文件)
|
||||
SingleFile bool `json:"single_file"` // 单文件模式(直接下载文件而非 sparse-checkout)
|
||||
Proxy string `json:"proxy"` // 代理类型: none, ghproxy, mirror, custom
|
||||
ProxyURL string `json:"proxy_url"` // 自定义代理地址
|
||||
AuthToken string `json:"auth_token"` // 认证 Token
|
||||
WhitelistPaths string `json:"whitelist_paths"` // 同步时保留的路径及脚本筛选白名单关键词,逗号或竖线分割
|
||||
Blacklist string `json:"blacklist"` // 脚本筛选黑名单关键词,竖线分割
|
||||
Dependence string `json:"dependence"` // 脚本依赖文件关键词,竖线分割
|
||||
Extensions string `json:"extensions"` // 脚本文件后缀关键词,竖线分割
|
||||
AutoAddCron bool `json:"auto_add_cron"` // 自动解析脚本注释添加定时任务
|
||||
CommentToTask string `json:"commenttotask"` // 兼容 QL 格式任务脚本注释解析
|
||||
RepoSource string `json:"repo_source"` // 仓库来源,如果是选择了这个 ql 导入的仓库,= ql
|
||||
RepoDirName string `json:"repo_dir_name"` // 自定义仓库目录名
|
||||
}
|
||||
|
||||
// TaskConfig 任务配置 RepoConfig+TaskConfig=task.config
|
||||
type TaskConfig struct {
|
||||
Concurrency int `json:"$task_concurrency"` // 0: disable concurrency, 1: enable concurrency
|
||||
AllEnvs bool `json:"$task_all_envs"` // 开启则注入全部环境变量
|
||||
}
|
||||
|
||||
// Task 代表一个计划任务
|
||||
type Task struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Name string `json:"name" gorm:"size:255;not null"`
|
||||
Remark string `json:"remark" gorm:"size:255;default:''"`
|
||||
PinType string `json:"pin_type" gorm:"size:20;default:none;index"` // 置顶类型: constant.PinTypeNone, constant.PinTypeTop
|
||||
Command BigText `json:"command"` // 普通任务的命令
|
||||
PreCommand BigText `json:"pre_command"` // 执行前的命令
|
||||
PostCommand BigText `json:"post_command"` // 执行后的命令
|
||||
Tags string `json:"tags" gorm:"-"` // 标签,逗号分隔
|
||||
Type string `json:"type" gorm:"size:20;default:'task'"` // 任务类型: constant.TaskTypeNormal, constant.TaskTypeRepo
|
||||
TriggerType string `json:"trigger_type" gorm:"size:25;default:'cron'"` // 触发类型: constant.TriggerTypeCron, constant.TriggerTypeTaskPoolStartup
|
||||
Config BigText `json:"config"` // 配置 JSON(仓库同步配置等)
|
||||
Schedule string `json:"schedule" gorm:"size:100"` // cron 表达式
|
||||
Timeout int `json:"timeout" gorm:"default:30"` // 超时时间(分钟),默认30分钟
|
||||
WorkDir string `json:"work_dir" gorm:"size:255;default:''"` // 工作目录,为空则使用 scripts 目录
|
||||
CleanConfig string `json:"clean_config" gorm:"size:255;default:''"` // 清理配置 JSON
|
||||
Envs BigText `json:"envs" gorm:"-"` // 环境变量ID列表,逗号分隔
|
||||
Languages TaskLanguages `json:"languages" gorm:"type:text"` // 针对本地任务的语言配置列表
|
||||
AgentID *string `json:"agent_id" gorm:"size:20;index"` // Agent ID,为空表示本地执行
|
||||
RetryCount int `json:"retry_count" gorm:"default:0"` // 失败重试次数
|
||||
RetryInterval int `json:"retry_interval" gorm:"default:0"` // 失败重试间隔(秒)
|
||||
RandomRange int `json:"random_range" gorm:"default:0"` // 随机延迟范围(秒)
|
||||
Enabled *bool `json:"enabled" gorm:"default:true"`
|
||||
RunningGo BigText `json:"running_go"` // 正在运行的 go routine id 数组 (JSON)
|
||||
RuntimeEnvs []string `json:"-" gorm:"-"` // 运行时环境变量(非持久化)
|
||||
RuntimeSecrets []string `json:"-" gorm:"-"` // 运行时安全机密(非持久化)
|
||||
LastRun *LocalTime `json:"last_run"`
|
||||
NextRun *LocalTime `json:"next_run"`
|
||||
SourceID string `json:"source_id" gorm:"size:255;index"` // 脚本资源唯一标识(路径 sanitized)
|
||||
RepoTaskID string `json:"repo_task_id" gorm:"size:20;index"` // 所属的仓库任务 ID
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (t *Task) IsRunning() bool {
|
||||
if string(t.RunningGo) == "" || string(t.RunningGo) == "[]" {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (Task) TableName() string {
|
||||
return constant.TablePrefix + "tasks"
|
||||
}
|
||||
|
||||
func (t *Task) GetID() string {
|
||||
return t.ID
|
||||
}
|
||||
|
||||
func (t *Task) GetName() string {
|
||||
return t.Name
|
||||
}
|
||||
|
||||
func (t *Task) GetCommand() string {
|
||||
return string(t.Command)
|
||||
}
|
||||
|
||||
func (t *Task) GetPreCommand() string {
|
||||
return string(t.PreCommand)
|
||||
}
|
||||
|
||||
func (t *Task) GetPostCommand() string {
|
||||
return string(t.PostCommand)
|
||||
}
|
||||
|
||||
func (t *Task) GetTimeout() int {
|
||||
return t.Timeout
|
||||
}
|
||||
|
||||
func (t *Task) GetWorkDir() string {
|
||||
return t.WorkDir
|
||||
}
|
||||
|
||||
func (t *Task) GetEnvs() string {
|
||||
return string(t.Envs)
|
||||
}
|
||||
|
||||
func (t *Task) GetLanguages() []map[string]string {
|
||||
return []map[string]string(t.Languages)
|
||||
}
|
||||
|
||||
func (t *Task) GetEnvVars() []string {
|
||||
return t.RuntimeEnvs
|
||||
}
|
||||
|
||||
func (t *Task) GetSecrets() []string {
|
||||
return t.RuntimeSecrets
|
||||
}
|
||||
|
||||
func (t *Task) GetUseMise() bool {
|
||||
return t.AgentID == nil || *t.AgentID == ""
|
||||
}
|
||||
|
||||
func (t *Task) UseMise() bool {
|
||||
return t.GetUseMise()
|
||||
}
|
||||
|
||||
// CronTask 计划任务接口
|
||||
func (t *Task) GetSchedule() string {
|
||||
return t.Schedule
|
||||
}
|
||||
|
||||
func (t *Task) GetRandomRange() int {
|
||||
return t.RandomRange
|
||||
}
|
||||
|
||||
// TaskLog 代表任务执行的日志记录
|
||||
type TaskLog struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
TaskID string `json:"task_id" gorm:"size:20;index"`
|
||||
AgentID *string `json:"agent_id" gorm:"size:20;index"` // Agent ID,为空表示本地执行
|
||||
Command BigText `json:"command"`
|
||||
Output BigText `json:"-"` // gzip+base64 压缩后的日志
|
||||
Error BigText `json:"error"` // 额外的系统错误信息
|
||||
Status string `json:"status" gorm:"size:20;index"` // success, failed
|
||||
Duration int64 `json:"duration"` // 执行耗时(毫秒)
|
||||
ExitCode int `json:"exit_code"`
|
||||
StartTime *LocalTime `json:"start_time"`
|
||||
EndTime *LocalTime `json:"end_time"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
}
|
||||
|
||||
func (TaskLog) TableName() string {
|
||||
return constant.TablePrefix + "task_logs"
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/systime"
|
||||
)
|
||||
|
||||
const TimeFormat = "2006-01-02 15:04:05"
|
||||
|
||||
// LocalTime 自定义时间类型,JSON 序列化为 "年-月-日 时:分:秒" 格式
|
||||
type LocalTime time.Time
|
||||
|
||||
func (t LocalTime) MarshalJSON() ([]byte, error) {
|
||||
tt := time.Time(t)
|
||||
if tt.IsZero() {
|
||||
return []byte("null"), nil
|
||||
}
|
||||
// 统一输出为东八区时间
|
||||
tt = systime.InCST(tt)
|
||||
return []byte(fmt.Sprintf(`"%s"`, tt.Format(TimeFormat))), nil
|
||||
}
|
||||
|
||||
func (t *LocalTime) UnmarshalJSON(data []byte) error {
|
||||
if string(data) == "null" || string(data) == `""` {
|
||||
return nil
|
||||
}
|
||||
// 去掉引号
|
||||
s := string(data)
|
||||
if len(s) >= 2 && s[0] == '"' && s[len(s)-1] == '"' {
|
||||
s = s[1 : len(s)-1]
|
||||
}
|
||||
tt, err := time.ParseInLocation(TimeFormat, s, time.Local)
|
||||
if err != nil {
|
||||
// 尝试解析 ISO 格式
|
||||
tt, err = time.Parse(time.RFC3339, s)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
*t = LocalTime(tt)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t LocalTime) Value() (driver.Value, error) {
|
||||
return time.Time(t), nil
|
||||
}
|
||||
|
||||
func (t *LocalTime) Scan(v interface{}) error {
|
||||
if v == nil {
|
||||
return nil
|
||||
}
|
||||
switch val := v.(type) {
|
||||
case time.Time:
|
||||
*t = LocalTime(val)
|
||||
case string:
|
||||
tt, err := time.ParseInLocation(TimeFormat, val, time.Local)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
*t = LocalTime(tt)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t LocalTime) Time() time.Time {
|
||||
return time.Time(t)
|
||||
}
|
||||
|
||||
func Now() LocalTime {
|
||||
return LocalTime(systime.Now())
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// User represents a system user
|
||||
type User struct {
|
||||
ID string `json:"id" gorm:"primaryKey;size:20"`
|
||||
Username string `json:"username" gorm:"size:100;uniqueIndex;not null"`
|
||||
Password string `json:"password" gorm:"size:255;not null"`
|
||||
Email string `json:"email" gorm:"size:255"`
|
||||
Role string `json:"role" gorm:"size:20;default:user"` // admin, user
|
||||
TokenVersion int `json:"-" gorm:"default:1"` // 用于 JWT 失效校验
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (User) TableName() string {
|
||||
return constant.TablePrefix + "users"
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
package vo
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
// AgentVO 代理视图对象
|
||||
type AgentVO struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Status string `json:"status"`
|
||||
LastSeen *models.LocalTime `json:"last_seen"`
|
||||
IP string `json:"ip"`
|
||||
Version string `json:"version"`
|
||||
BuildTime string `json:"build_time"`
|
||||
Hostname string `json:"hostname"`
|
||||
OS string `json:"os"`
|
||||
Arch string `json:"arch"`
|
||||
ForceUpdate bool `json:"force_update"`
|
||||
Enabled bool `json:"enabled"`
|
||||
SchedulerConfig *AgentSchedulerConfigVO `json:"scheduler_config"`
|
||||
CreatedAt models.LocalTime `json:"created_at"`
|
||||
UpdatedAt models.LocalTime `json:"updated_at"`
|
||||
// 隐藏 Token 和 MachineID
|
||||
}
|
||||
|
||||
// ToAgentVO 将 Agent 模型转换为 AgentVO
|
||||
func ToAgentVO(agent *models.Agent) *AgentVO {
|
||||
if agent == nil {
|
||||
return nil
|
||||
}
|
||||
var schedulerConfigVO *AgentSchedulerConfigVO
|
||||
if agent.SchedulerConfig.WorkerCount > 0 {
|
||||
schedulerConfigVO = &AgentSchedulerConfigVO{
|
||||
WorkerCount: agent.SchedulerConfig.WorkerCount,
|
||||
QueueSize: agent.SchedulerConfig.QueueSize,
|
||||
RateInterval: int(agent.SchedulerConfig.RateInterval / time.Millisecond),
|
||||
Verbose: agent.SchedulerConfig.Verbose,
|
||||
StrictQueue: agent.SchedulerConfig.StrictQueue,
|
||||
}
|
||||
}
|
||||
return &AgentVO{
|
||||
ID: agent.ID,
|
||||
Name: agent.Name,
|
||||
Description: agent.Description,
|
||||
Status: agent.Status,
|
||||
LastSeen: agent.LastSeen,
|
||||
IP: agent.IP,
|
||||
Version: agent.Version,
|
||||
BuildTime: agent.BuildTime,
|
||||
Hostname: agent.Hostname,
|
||||
OS: agent.OS,
|
||||
Arch: agent.Arch,
|
||||
ForceUpdate: agent.ForceUpdate,
|
||||
Enabled: utils.DerefBool(agent.Enabled, true),
|
||||
SchedulerConfig: schedulerConfigVO,
|
||||
CreatedAt: agent.CreatedAt,
|
||||
UpdatedAt: agent.UpdatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
// ToAgentVOList 将 Agent 模型列表转换为 AgentVO 列表
|
||||
func ToAgentVOList(agents []*models.Agent) []*AgentVO {
|
||||
if agents == nil {
|
||||
return nil
|
||||
}
|
||||
vos := make([]*AgentVO, len(agents))
|
||||
for i, a := range agents {
|
||||
vos[i] = ToAgentVO(a)
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// ToAgentVOListFromModels 将 Agent 模型列表转换为 AgentVO 列表
|
||||
func ToAgentVOListFromModels(agents []models.Agent) []*AgentVO {
|
||||
vos := make([]*AgentVO, len(agents))
|
||||
for i := range agents {
|
||||
vos[i] = ToAgentVO(&agents[i])
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// AgentTokenVO 代理令牌视图对象
|
||||
type AgentTokenVO struct {
|
||||
ID string `json:"id"`
|
||||
Token string `json:"token"`
|
||||
Remark string `json:"remark"`
|
||||
MaxUses int `json:"max_uses"`
|
||||
UsedCount int `json:"used_count"`
|
||||
ExpiresAt *models.LocalTime `json:"expires_at"`
|
||||
Enabled bool `json:"enabled"`
|
||||
CreatedAt models.LocalTime `json:"created_at"`
|
||||
}
|
||||
|
||||
// ToAgentTokenVO 将 AgentToken 模型转换为 AgentTokenVO
|
||||
func ToAgentTokenVO(token *models.AgentToken) *AgentTokenVO {
|
||||
if token == nil {
|
||||
return nil
|
||||
}
|
||||
return &AgentTokenVO{
|
||||
ID: token.ID,
|
||||
Token: token.Token,
|
||||
Remark: token.Remark,
|
||||
MaxUses: token.MaxUses,
|
||||
UsedCount: token.UsedCount,
|
||||
ExpiresAt: token.ExpiresAt,
|
||||
Enabled: utils.DerefBool(token.Enabled, true),
|
||||
CreatedAt: token.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
// ToAgentTokenVOList 将 AgentToken 模型列表转换为 AgentTokenVO 列表
|
||||
func ToAgentTokenVOList(tokens []*models.AgentToken) []*AgentTokenVO {
|
||||
if tokens == nil {
|
||||
return nil
|
||||
}
|
||||
vos := make([]*AgentTokenVO, len(tokens))
|
||||
for i, t := range tokens {
|
||||
vos[i] = ToAgentTokenVO(t)
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// ToAgentTokenVOListFromModels 将 AgentToken 模型列表转换为 AgentTokenVO 列表
|
||||
func ToAgentTokenVOListFromModels(tokens []models.AgentToken) []*AgentTokenVO {
|
||||
vos := make([]*AgentTokenVO, len(tokens))
|
||||
for i := range tokens {
|
||||
vos[i] = ToAgentTokenVO(&tokens[i])
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// AgentSchedulerConfigVO 调度配置视图对象
|
||||
type AgentSchedulerConfigVO struct {
|
||||
WorkerCount int `json:"worker_count"`
|
||||
QueueSize int `json:"queue_size"`
|
||||
RateInterval int `json:"rate_interval"` // 毫秒
|
||||
Verbose bool `json:"verbose"`
|
||||
StrictQueue bool `json:"strict_queue"`
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package vo
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
// DependencyVO 依赖包视图对象
|
||||
type DependencyVO struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
Language string `json:"language"`
|
||||
LangVersion string `json:"lang_version"`
|
||||
Remark string `json:"remark"`
|
||||
Log string `json:"log,omitempty"` // 仅在需要时返回
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ToDependencyVO 将 Dependency 模型转换为 DependencyVO
|
||||
func ToDependencyVO(dep *models.Dependency) *DependencyVO {
|
||||
if dep == nil {
|
||||
return nil
|
||||
}
|
||||
return &DependencyVO{
|
||||
ID: dep.ID,
|
||||
Name: dep.Name,
|
||||
Version: dep.Version,
|
||||
Language: dep.Language,
|
||||
LangVersion: dep.LangVersion,
|
||||
Remark: dep.Remark,
|
||||
Log: string(dep.Log),
|
||||
CreatedAt: dep.CreatedAt,
|
||||
UpdatedAt: dep.UpdatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
// ToDependencyVOListFromModels 将 Dependency 模型列表转换为 DependencyVO 列表
|
||||
func ToDependencyVOListFromModels(deps []models.Dependency) []*DependencyVO {
|
||||
vos := make([]*DependencyVO, len(deps))
|
||||
for i := range deps {
|
||||
vos[i] = ToDependencyVO(&deps[i])
|
||||
}
|
||||
return vos
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package vo
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
)
|
||||
|
||||
// ScriptVO 脚本视图对象
|
||||
type ScriptVO struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Content string `json:"content,omitempty"` // 仅在拉取详情时返回
|
||||
CreatedAt models.LocalTime `json:"created_at"`
|
||||
UpdatedAt models.LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ToScriptVO 将 Script 模型转换为 ScriptVO
|
||||
func ToScriptVO(script *models.Script) *ScriptVO {
|
||||
if script == nil {
|
||||
return nil
|
||||
}
|
||||
return &ScriptVO{
|
||||
ID: script.ID,
|
||||
Name: script.Name,
|
||||
Content: string(script.Content),
|
||||
CreatedAt: script.CreatedAt,
|
||||
UpdatedAt: script.UpdatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
// ToScriptVOListFromModels 将 Script 模型列表转换为 ScriptVO 列表
|
||||
func ToScriptVOListFromModels(scripts []models.Script) []*ScriptVO {
|
||||
vos := make([]*ScriptVO, len(scripts))
|
||||
for i := range scripts {
|
||||
vos[i] = ToScriptVO(&scripts[i])
|
||||
}
|
||||
return vos
|
||||
}
|
||||
@@ -0,0 +1,108 @@
|
||||
package vo
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
// UserVO 用户视图对象
|
||||
type UserVO struct {
|
||||
ID string `json:"id"`
|
||||
Username string `json:"username"`
|
||||
Email string `json:"email"`
|
||||
Role string `json:"role"`
|
||||
CreatedAt models.LocalTime `json:"created_at"`
|
||||
UpdatedAt models.LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ToUserVO 将 User 模型转换为 UserVO
|
||||
func ToUserVO(user *models.User) *UserVO {
|
||||
if user == nil {
|
||||
return nil
|
||||
}
|
||||
return &UserVO{
|
||||
ID: user.ID,
|
||||
Username: user.Username,
|
||||
Email: user.Email,
|
||||
Role: user.Role,
|
||||
CreatedAt: user.CreatedAt,
|
||||
UpdatedAt: user.UpdatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
// EnvVO 环境变量视图对象
|
||||
type EnvVO struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Value string `json:"value"`
|
||||
Remark string `json:"remark"`
|
||||
Type string `json:"type"`
|
||||
Tags string `json:"tags"`
|
||||
Hidden bool `json:"hidden"`
|
||||
Enabled bool `json:"enabled"`
|
||||
CreatedAt models.LocalTime `json:"created_at"`
|
||||
UpdatedAt models.LocalTime `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ToEnvVO 将 Env 模型转换为 EnvVO
|
||||
func ToEnvVO(env *models.EnvironmentVariable) *EnvVO {
|
||||
if env == nil {
|
||||
return nil
|
||||
}
|
||||
val := string(env.Value)
|
||||
if env.Type == constant.EnvTypeSecret {
|
||||
val = "********"
|
||||
}
|
||||
return &EnvVO{
|
||||
ID: env.ID,
|
||||
Name: env.Name,
|
||||
Value: val,
|
||||
Remark: env.Remark,
|
||||
Type: env.Type,
|
||||
Tags: env.Tags,
|
||||
Hidden: utils.DerefBool(env.Hidden, true),
|
||||
Enabled: utils.DerefBool(env.Enabled, true),
|
||||
CreatedAt: env.CreatedAt,
|
||||
UpdatedAt: env.UpdatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
// ToEnvVOList 将 Env 模型列表转换为 EnvVO 列表
|
||||
func ToEnvVOList(envs []*models.EnvironmentVariable) []*EnvVO {
|
||||
if envs == nil {
|
||||
return nil
|
||||
}
|
||||
vos := make([]*EnvVO, len(envs))
|
||||
for i, e := range envs {
|
||||
vos[i] = ToEnvVO(e)
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// ToEnvVOListFromModels 将 Env 模型列表转换为 EnvVO 列表
|
||||
func ToEnvVOListFromModels(envs []models.EnvironmentVariable) []*EnvVO {
|
||||
vos := make([]*EnvVO, len(envs))
|
||||
for i := range envs {
|
||||
vos[i] = ToEnvVO(&envs[i])
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// LoginLogVO 登录日志视图对象
|
||||
type LoginLogVO struct {
|
||||
ID string `json:"id"`
|
||||
Username string `json:"username"`
|
||||
IP string `json:"ip"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
Status string `json:"status"`
|
||||
Message string `json:"message"`
|
||||
CreatedAt models.LocalTime `json:"created_at"`
|
||||
}
|
||||
|
||||
// TokenConfig Token 配置结构体
|
||||
type TokenConfig struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
Token string `json:"token"`
|
||||
ExpireAt string `json:"expire_at"`
|
||||
}
|
||||
@@ -0,0 +1,263 @@
|
||||
package vo
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/executor"
|
||||
"github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/utils"
|
||||
)
|
||||
|
||||
// TaskCreateReq 任务创建请求
|
||||
type TaskCreateReq struct {
|
||||
Name string `json:"name" binding:"required" example:"测试任务"`
|
||||
Remark string `json:"remark" example:"备注信息"`
|
||||
Command string `json:"command" example:"echo 'Hello World'"`
|
||||
PreCommand string `json:"pre_command" example:"echo 'pre'"`
|
||||
PostCommand string `json:"post_command" example:"echo 'post'"`
|
||||
Tags string `json:"tags" example:"test,dev"`
|
||||
Type string `json:"type" example:"repo"` // 可以是 common, repo 等
|
||||
Config string `json:"config" swaggertype:"string" example:"{\"source_url\":\"https://github.com/abc/repo\",\"branch\":\"main\"}"`
|
||||
Schedule string `json:"schedule" example:"0 0 * * *"`
|
||||
Timeout int `json:"timeout" example:"3600"`
|
||||
WorkDir string `json:"work_dir" example:"/tmp"`
|
||||
CleanConfig string `json:"clean_config" example:"true"`
|
||||
Envs string `json:"envs" example:"{\"ENV_VAR\":\"value\"}"`
|
||||
Languages models.TaskLanguages `json:"languages"`
|
||||
AgentID *string `json:"agent_id" example:"agent-1"`
|
||||
TriggerType string `json:"trigger_type" example:"cron"`
|
||||
RetryCount int `json:"retry_count" example:"3"`
|
||||
RetryInterval int `json:"retry_interval" example:"60"`
|
||||
RandomRange int `json:"random_range" example:"10"`
|
||||
PinType string `json:"pin_type" example:"time"`
|
||||
}
|
||||
|
||||
// TaskUpdateReq 任务更新请求
|
||||
type TaskUpdateReq struct {
|
||||
Name string `json:"name" example:"测试任务"`
|
||||
Remark string `json:"remark" example:"备注信息"`
|
||||
Command string `json:"command" example:"echo 'Hello World'"`
|
||||
PreCommand string `json:"pre_command" example:"echo 'pre'"`
|
||||
PostCommand string `json:"post_command" example:"echo 'post'"`
|
||||
Tags string `json:"tags" example:"test,dev"`
|
||||
Type string `json:"type" example:"repo"`
|
||||
Config string `json:"config" swaggertype:"string" example:"{\"source_url\":\"https://github.com/abc/repo\",\"branch\":\"main\"}"`
|
||||
Schedule string `json:"schedule" example:"0 0 * * *"`
|
||||
Timeout int `json:"timeout" example:"3600"`
|
||||
WorkDir string `json:"work_dir" example:"/tmp"`
|
||||
CleanConfig string `json:"clean_config" example:"true"`
|
||||
Envs string `json:"envs" example:"{\"ENV_VAR\":\"value\"}"`
|
||||
Enabled bool `json:"enabled" example:"true"`
|
||||
Languages models.TaskLanguages `json:"languages"`
|
||||
AgentID *string `json:"agent_id" example:"agent-1"`
|
||||
TriggerType string `json:"trigger_type" example:"cron"`
|
||||
RetryCount int `json:"retry_count" example:"3"`
|
||||
RetryInterval int `json:"retry_interval" example:"60"`
|
||||
RandomRange int `json:"random_range" example:"10"`
|
||||
PinType string `json:"pin_type" example:"time"`
|
||||
}
|
||||
|
||||
// TaskVO 任务视图对象
|
||||
type TaskVO struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Remark string `json:"remark"`
|
||||
Command string `json:"command"`
|
||||
PreCommand string `json:"pre_command"`
|
||||
PostCommand string `json:"post_command"`
|
||||
Tags string `json:"tags"`
|
||||
Type string `json:"type"`
|
||||
TriggerType string `json:"trigger_type"`
|
||||
Config string `json:"config"`
|
||||
Schedule string `json:"schedule"`
|
||||
Timeout int `json:"timeout"`
|
||||
WorkDir string `json:"work_dir"`
|
||||
CleanConfig string `json:"clean_config"`
|
||||
Envs string `json:"envs"`
|
||||
Languages models.TaskLanguages `json:"languages"`
|
||||
AgentID *string `json:"agent_id"`
|
||||
RepoTaskID string `json:"repo_task_id"`
|
||||
Enabled bool `json:"enabled"`
|
||||
RetryCount int `json:"retry_count"`
|
||||
RetryInterval int `json:"retry_interval"`
|
||||
RandomRange int `json:"random_range"`
|
||||
PinType string `json:"pin_type"`
|
||||
LastRun *models.LocalTime `json:"last_run"`
|
||||
NextRun *models.LocalTime `json:"next_run"`
|
||||
CreatedAt models.LocalTime `json:"created_at"`
|
||||
UpdatedAt models.LocalTime `json:"updated_at"`
|
||||
RunningStatus string `json:"running_status"`
|
||||
}
|
||||
|
||||
// ToTaskVO 将 Task 模型转换为 TaskVO
|
||||
func ToTaskVO(task *models.Task) *TaskVO {
|
||||
if task == nil {
|
||||
return nil
|
||||
}
|
||||
return &TaskVO{
|
||||
ID: task.ID,
|
||||
Name: task.Name,
|
||||
Remark: task.Remark,
|
||||
Command: string(task.Command),
|
||||
PreCommand: string(task.PreCommand),
|
||||
PostCommand: string(task.PostCommand),
|
||||
Tags: task.Tags,
|
||||
Type: task.Type,
|
||||
TriggerType: task.TriggerType,
|
||||
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,
|
||||
RepoTaskID: task.RepoTaskID,
|
||||
Enabled: utils.DerefBool(task.Enabled, true),
|
||||
RetryCount: task.RetryCount,
|
||||
RetryInterval: task.RetryInterval,
|
||||
RandomRange: task.RandomRange,
|
||||
PinType: task.PinType,
|
||||
LastRun: task.LastRun,
|
||||
NextRun: task.NextRun,
|
||||
CreatedAt: task.CreatedAt,
|
||||
UpdatedAt: task.UpdatedAt,
|
||||
RunningStatus: func() string {
|
||||
if task.IsRunning() {
|
||||
return "running"
|
||||
}
|
||||
return "idle"
|
||||
}(),
|
||||
}
|
||||
}
|
||||
|
||||
// ToTaskVOList 将 Task 模型列表转换为 TaskVO 列表
|
||||
func ToTaskVOList(tasks []*models.Task) []*TaskVO {
|
||||
if tasks == nil {
|
||||
return nil
|
||||
}
|
||||
vos := make([]*TaskVO, len(tasks))
|
||||
for i, t := range tasks {
|
||||
vos[i] = ToTaskVO(t)
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// ToTaskVOListFromModels 将 Task 模型列表转换为 TaskVO 列表
|
||||
func ToTaskVOListFromModels(tasks []models.Task) []*TaskVO {
|
||||
vos := make([]*TaskVO, len(tasks))
|
||||
for i := range tasks {
|
||||
vos[i] = ToTaskVO(&tasks[i])
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// TaskLogVO 任务历史视图对象
|
||||
type TaskLogVO struct {
|
||||
ID string `json:"id"`
|
||||
TaskID string `json:"task_id"`
|
||||
TaskName string `json:"task_name"`
|
||||
TaskType string `json:"task_type"`
|
||||
AgentID *string `json:"agent_id"`
|
||||
Command string `json:"command"`
|
||||
Error string `json:"error"`
|
||||
Status string `json:"status"`
|
||||
Duration int64 `json:"duration"`
|
||||
ExitCode int `json:"exit_code"`
|
||||
StartTime *models.LocalTime `json:"start_time"`
|
||||
EndTime *models.LocalTime `json:"end_time"`
|
||||
CreatedAt models.LocalTime `json:"created_at"`
|
||||
Output string `json:"output,omitempty"`
|
||||
}
|
||||
|
||||
// ToTaskLogVO 将 TaskLog 模型转换为 TaskLogVO
|
||||
// Note: This function assumes the Task field within models.TaskLog is preloaded
|
||||
// or that taskName and taskType are provided from an external source.
|
||||
func ToTaskLogVO(log *models.TaskLog) *TaskLogVO {
|
||||
if log == nil {
|
||||
return nil
|
||||
}
|
||||
return &TaskLogVO{
|
||||
ID: log.ID,
|
||||
TaskID: log.TaskID,
|
||||
AgentID: log.AgentID,
|
||||
Command: string(log.Command),
|
||||
Error: string(log.Error),
|
||||
Status: log.Status,
|
||||
Duration: log.Duration,
|
||||
ExitCode: log.ExitCode,
|
||||
StartTime: log.StartTime,
|
||||
EndTime: log.EndTime,
|
||||
CreatedAt: log.CreatedAt,
|
||||
Output: string(log.Output),
|
||||
}
|
||||
}
|
||||
|
||||
// ToTaskLogVOList 将 TaskLog 模型列表转换为 TaskLogVO 列表
|
||||
func ToTaskLogVOList(logs []*models.TaskLog) []*TaskLogVO {
|
||||
if logs == nil {
|
||||
return nil
|
||||
}
|
||||
vos := make([]*TaskLogVO, len(logs))
|
||||
for i, l := range logs {
|
||||
vos[i] = ToTaskLogVO(l)
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// ToTaskLogVOListFromModels 将 TaskLog 模型列表转换为 TaskLogVO 列表
|
||||
func ToTaskLogVOListFromModels(logs []models.TaskLog) []*TaskLogVO {
|
||||
vos := make([]*TaskLogVO, len(logs))
|
||||
for i := range logs {
|
||||
vos[i] = ToTaskLogVO(&logs[i])
|
||||
}
|
||||
return vos
|
||||
}
|
||||
|
||||
// ExecutionResultVO 任务执行结果视图对象
|
||||
type ExecutionResultVO struct {
|
||||
TaskID string `json:"task_id"`
|
||||
LogID string `json:"log_id,omitempty"`
|
||||
Success bool `json:"success"`
|
||||
Status string `json:"status"`
|
||||
Output string `json:"output,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
Duration int64 `json:"duration,omitempty"`
|
||||
ExitCode int `json:"exit_code,omitempty"`
|
||||
StartTime string `json:"start_time,omitempty"`
|
||||
EndTime string `json:"end_time,omitempty"`
|
||||
}
|
||||
|
||||
// ToExecutionResultVO 将 ExecutionResult 转换为 ExecutionResultVO
|
||||
func ToExecutionResultVO(res *executor.ExecutionResult) *ExecutionResultVO {
|
||||
if res == nil {
|
||||
return nil
|
||||
}
|
||||
vo := &ExecutionResultVO{
|
||||
TaskID: res.TaskID,
|
||||
LogID: res.LogID,
|
||||
Success: res.Success,
|
||||
Status: res.Status,
|
||||
Output: res.Output,
|
||||
Error: res.Error,
|
||||
Duration: res.Duration,
|
||||
ExitCode: res.ExitCode,
|
||||
}
|
||||
if !res.StartTime.IsZero() {
|
||||
vo.StartTime = res.StartTime.Format("2006-01-02 15:04:05")
|
||||
}
|
||||
if !res.EndTime.IsZero() {
|
||||
vo.EndTime = res.EndTime.Format("2006-01-02 15:04:05")
|
||||
}
|
||||
return vo
|
||||
}
|
||||
|
||||
// ToExecutionResultVOList 将 ExecutionResult 列表转换为 ExecutionResultVO 列表
|
||||
func ToExecutionResultVOList(results []executor.ExecutionResult) []*ExecutionResultVO {
|
||||
if results == nil {
|
||||
return nil
|
||||
}
|
||||
vos := make([]*ExecutionResultVO, len(results))
|
||||
for i := range results {
|
||||
vos[i] = ToExecutionResultVO(&results[i])
|
||||
}
|
||||
return vos
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
package vo
|
||||
|
||||
// WSMessage 通用 WebSocket 消息结构
|
||||
type WSMessage struct {
|
||||
Type string `json:"type"` // 事件类型: task_status, notice, system_stats
|
||||
Timestamp int64 `json:"timestamp"` // 毫秒时间戳
|
||||
Payload interface{} `json:"payload"` // 负载数据
|
||||
}
|
||||
@@ -0,0 +1,350 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/middleware"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func initPublicAPIRoutes(api *gin.RouterGroup, c *Controllers) {
|
||||
// Health check (无需认证)
|
||||
api.GET("/ping", func(ctx *gin.Context) {
|
||||
ctx.JSON(200, gin.H{"message": "pong"})
|
||||
})
|
||||
|
||||
// Install routes (无需认证,仅在未安装时可用)
|
||||
install := api.Group("/install")
|
||||
{
|
||||
install.GET("/status", c.Install.GetInstallStatus)
|
||||
install.POST("", c.Install.Install)
|
||||
}
|
||||
|
||||
// api.GET("/debug/goroutines", func(ctx *gin.Context) {
|
||||
// buf := make([]byte, 1024*1024)
|
||||
// n := runtime.Stack(buf, true)
|
||||
// ctx.Data(200, "text/plain; charset=utf-8", buf[:n])
|
||||
// })
|
||||
|
||||
// Authentication routes (无需认证)
|
||||
auth := api.Group("/auth")
|
||||
{
|
||||
auth.POST("/login", c.Auth.Login)
|
||||
auth.POST("/logout", c.Auth.Logout)
|
||||
// auth.POST("/register", c.Auth.Register)
|
||||
}
|
||||
|
||||
// 公开的站点设置(无需认证)
|
||||
api.GET("/settings/public", c.Settings.GetPublicSiteSettings)
|
||||
|
||||
// 隧道模式 (被控端反向连入,使用独立 Token 做 WebSocket 鉴权)
|
||||
api.GET("/interconnect/tunnel", c.Interconnect.HandleTunnel)
|
||||
// 子节点主动上报监控数据 (无中间件鉴权,内部鉴权)
|
||||
api.POST("/interconnect/report", c.Interconnect.ReportMonitorData)
|
||||
|
||||
// 内部使用的 API(仅限本地调用,无需 Bearer 认证)
|
||||
internalAPI := api.Group("/internal")
|
||||
internalAPI.Use(middleware.LocalhostOnly())
|
||||
{
|
||||
internalAPI.POST("/tasks/sync-repo-status", c.Task.SyncRepoTasks)
|
||||
internalAPI.POST("/tasks/execute/:id", c.Executor.ExecuteTask)
|
||||
internalAPI.POST("/tasks/toggle/:id", c.Task.ToggleTask)
|
||||
}
|
||||
}
|
||||
|
||||
func initAuthorizedAPIRoutes(api *gin.RouterGroup, c *Controllers) {
|
||||
authorized := api.Group("")
|
||||
authorized.Use(middleware.AuthRequired())
|
||||
{
|
||||
// 获取当前用户 (普通用户即可访问)
|
||||
authorized.GET("/auth/me", c.Auth.GetCurrentUser)
|
||||
|
||||
// 以下管理接口需要管理员权限
|
||||
adminOnly := authorized.Group("")
|
||||
adminOnly.Use(middleware.AdminRequired())
|
||||
{
|
||||
registerDashboardRoutes(adminOnly, c)
|
||||
registerTaskRoutes(adminOnly, c)
|
||||
registerEnvRoutes(adminOnly, c)
|
||||
registerScriptRoutes(adminOnly, c)
|
||||
registerFileRoutes(adminOnly, c)
|
||||
registerLogRoutes(adminOnly, c)
|
||||
registerTerminalRoutes(adminOnly, c)
|
||||
registerSettingsRoutes(adminOnly, c)
|
||||
registerDependencyRoutes(adminOnly, c)
|
||||
registerAgentRoutes(adminOnly, c)
|
||||
registerMiseRoutes(adminOnly, c)
|
||||
registerNotificationRoutes(adminOnly, c)
|
||||
registerAppLogRoutes(adminOnly, c)
|
||||
registerSystemWSRoutes(adminOnly, c)
|
||||
registerWebUIRoutes(adminOnly, c)
|
||||
registerMonitorRoutes(adminOnly, c)
|
||||
registerInterconnectRoutes(adminOnly, c)
|
||||
registerSystemRoutes(adminOnly, c)
|
||||
}
|
||||
}
|
||||
|
||||
// 通知发送 API(使用通知 Token 认证,供脚本调用)
|
||||
notifyAPI := api.Group("/notify")
|
||||
notifyAPI.Use(middleware.NotifyTokenAuth())
|
||||
{
|
||||
notifyAPI.POST("/send", c.Notification.SendNotification)
|
||||
}
|
||||
}
|
||||
|
||||
func registerDashboardRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
g.GET("/stats", c.Dashboard.GetStats)
|
||||
g.GET("/sentence", c.Dashboard.GetSentence)
|
||||
g.GET("/sendstats", c.Dashboard.GetSendStats)
|
||||
g.GET("/taskstats", c.Dashboard.GetTaskStats)
|
||||
}
|
||||
|
||||
func registerTaskRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
tasks := g.Group("/tasks")
|
||||
{
|
||||
tasks.POST("", c.Task.CreateTask)
|
||||
tasks.GET("", c.Task.GetTasks)
|
||||
tasks.GET("/:id", c.Task.GetTask)
|
||||
tasks.POST("/bulk_save", c.Task.BulkSaveTask)
|
||||
tasks.PUT("/:id", c.Task.UpdateTask)
|
||||
tasks.DELETE("/:id", c.Task.DeleteTask)
|
||||
tasks.POST("/batch-delete", c.Task.BatchDeleteTasks)
|
||||
tasks.DELETE("/batch-by-query", c.Task.BatchDeleteByQuery)
|
||||
tasks.POST("/stop/:logID", c.Task.StopTask)
|
||||
tasks.GET("/tags", c.Task.GetTags)
|
||||
}
|
||||
|
||||
execution := g.Group("/execute")
|
||||
{
|
||||
execution.POST("/task/:id", c.Executor.ExecuteTask)
|
||||
execution.POST("/command", c.Executor.ExecuteCommand)
|
||||
execution.GET("/results", c.Executor.GetLastResults)
|
||||
}
|
||||
}
|
||||
|
||||
func registerEnvRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
env := g.Group("/env")
|
||||
{
|
||||
env.GET("/secret-status", c.Env.GetSecretStatus)
|
||||
env.GET("/tags", c.Env.GetTags)
|
||||
env.POST("", c.Env.CreateEnvVar)
|
||||
env.POST("/bulk_save", c.Env.BulkSaveEnv)
|
||||
env.GET("", c.Env.GetEnvVars)
|
||||
env.GET("/all", c.Env.GetAllEnvVars)
|
||||
env.GET("/:id", c.Env.GetEnvVar)
|
||||
env.GET("/:id/tasks", c.Env.GetAssociatedTasks)
|
||||
env.PUT("/:id", c.Env.UpdateEnvVar)
|
||||
env.DELETE("/:id", c.Env.DeleteEnvVar)
|
||||
}
|
||||
}
|
||||
|
||||
func registerScriptRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
scripts := g.Group("/scripts")
|
||||
{
|
||||
scripts.POST("", c.Script.CreateScript)
|
||||
scripts.GET("", c.Script.GetScripts)
|
||||
scripts.GET("/:id", c.Script.GetScript)
|
||||
scripts.PUT("/:id", c.Script.UpdateScript)
|
||||
scripts.DELETE("/:id", c.Script.DeleteScript)
|
||||
}
|
||||
}
|
||||
|
||||
func registerFileRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
files := g.Group("/files")
|
||||
{
|
||||
files.GET("/tree", c.File.GetFileTree)
|
||||
files.GET("/content", c.File.GetFileContent)
|
||||
files.GET("/download", c.File.DownloadFile)
|
||||
files.GET("/download-zip", c.File.DownloadZip)
|
||||
files.POST("/content", c.File.SaveFileContent)
|
||||
files.POST("/create", c.File.CreateFile)
|
||||
files.POST("/delete", c.File.DeleteFile)
|
||||
files.POST("/rename", c.File.RenameFile)
|
||||
files.POST("/move", c.File.MoveFile)
|
||||
files.POST("/copy", c.File.CopyFile)
|
||||
files.POST("/upload", c.File.UploadArchive)
|
||||
files.POST("/uploadfiles", c.File.UploadFiles)
|
||||
}
|
||||
}
|
||||
|
||||
func registerLogRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
logs := g.Group("/logs")
|
||||
{
|
||||
logs.GET("", c.Log.GetLogs)
|
||||
logs.POST("/clear", c.Log.ClearLogs)
|
||||
logs.GET("/sse", c.LogSSE.StreamLog)
|
||||
logs.GET("/:id", c.Log.GetLogDetail)
|
||||
logs.DELETE("/:id", c.Log.DeleteLog)
|
||||
}
|
||||
}
|
||||
|
||||
func registerTerminalRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
g.GET("/terminal/ws", c.Terminal.HandleWebSocket)
|
||||
// g.POST("/terminal/exec", c.Terminal.ExecuteShellCommand) // 暂未使用,已注释
|
||||
g.GET("/terminal/cmds", c.Terminal.GetCommands)
|
||||
}
|
||||
|
||||
func registerSettingsRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
settings := g.Group("/settings")
|
||||
{
|
||||
settings.POST("/password", c.Settings.ChangePassword)
|
||||
settings.GET("/site", c.Settings.GetSiteSettings)
|
||||
settings.PUT("/site", c.Settings.UpdateSiteSettings)
|
||||
settings.POST("/site/openapi-token/generate", c.Settings.GenerateOpenapiToken)
|
||||
settings.GET("/paths", c.Settings.GetPaths)
|
||||
settings.GET("/scheduler", c.Settings.GetSchedulerSettings)
|
||||
settings.PUT("/scheduler", c.Settings.UpdateSchedulerSettings)
|
||||
settings.GET("/about", c.Settings.GetAbout)
|
||||
settings.GET("/changelog", c.Settings.GetChangelog)
|
||||
settings.GET("/loginlogs", c.Settings.GetLoginLogs)
|
||||
settings.POST("/backup", c.Settings.CreateBackup)
|
||||
settings.GET("/backup/status", c.Settings.GetBackupStatus)
|
||||
settings.GET("/backup/download", c.Settings.DownloadBackup)
|
||||
settings.POST("/restore", c.Settings.RestoreBackup)
|
||||
// 通用设置接口
|
||||
settings.GET("/:section", c.Settings.GetSectionSettings)
|
||||
settings.PUT("/:section", c.Settings.UpdateSectionSettings)
|
||||
settings.GET("/:section/:key", c.Settings.GetSetting)
|
||||
settings.POST("/:section/:key/generate", c.Settings.GenerateSettingToken)
|
||||
}
|
||||
}
|
||||
|
||||
func registerDependencyRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
deps := g.Group("/deps")
|
||||
{
|
||||
deps.GET("", c.Dependency.List)
|
||||
deps.POST("", c.Dependency.Create)
|
||||
deps.DELETE("/:id", c.Dependency.Delete)
|
||||
deps.POST("/install", c.Dependency.Install)
|
||||
deps.POST("/install-cmd", c.Dependency.GetInstallCommand)
|
||||
deps.POST("/uninstall/:id", c.Dependency.Uninstall)
|
||||
deps.POST("/reinstall/:id", c.Dependency.Reinstall)
|
||||
deps.POST("/reinstall-all", c.Dependency.ReinstallAll)
|
||||
deps.POST("/reinstall-all-cmd", c.Dependency.GetReinstallAllCommand)
|
||||
deps.POST("/batch-install-cmd", c.Dependency.GetBatchInstallCommand)
|
||||
deps.POST("/import", c.Dependency.ParseAndImport)
|
||||
deps.GET("/installed", c.Dependency.GetInstalled)
|
||||
deps.GET("/install-suggest-cmd", c.Dependency.GetDepInstallCommand)
|
||||
}
|
||||
}
|
||||
|
||||
func registerAgentRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
agents := g.Group("/agents")
|
||||
{
|
||||
agents.GET("", c.Agent.List)
|
||||
agents.GET("/version", c.Agent.GetVersion)
|
||||
agents.PUT("/:id", c.Agent.Update)
|
||||
agents.DELETE("/:id", c.Agent.Delete)
|
||||
agents.POST("/:id/token", c.Agent.RegenerateToken)
|
||||
agents.POST("/:id/update", c.Agent.ForceUpdate)
|
||||
// 令牌管理
|
||||
agents.GET("/tokens", c.Agent.ListTokens)
|
||||
agents.POST("/tokens", c.Agent.CreateToken)
|
||||
agents.DELETE("/tokens/:id", c.Agent.DeleteToken)
|
||||
}
|
||||
|
||||
// Agent API(供前端调用,保持在 v1 下)
|
||||
agentAPIv1 := g.Group("/agent")
|
||||
{
|
||||
agentAPIv1.GET("/download", c.Agent.Download)
|
||||
}
|
||||
}
|
||||
|
||||
func registerMiseRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
mise := g.Group("/mise")
|
||||
{
|
||||
mise.GET("/ls", c.Mise.List)
|
||||
mise.POST("/sync", c.Mise.Sync)
|
||||
mise.GET("/plugins", c.Mise.Plugins)
|
||||
mise.GET("/versions", c.Mise.Versions)
|
||||
mise.GET("/verify-cmd", c.Mise.VerifyCommand)
|
||||
mise.POST("/use-global", c.Mise.UseGlobal)
|
||||
mise.POST("/unset-global", c.Mise.UnsetGlobal)
|
||||
mise.GET("/envs", c.Mise.Envs)
|
||||
mise.POST("/envs", c.Mise.SetEnv)
|
||||
mise.DELETE("/envs", c.Mise.UnsetEnv)
|
||||
}
|
||||
}
|
||||
|
||||
func registerNotificationRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
notify := g.Group("/notify")
|
||||
{
|
||||
notify.GET("/types", c.Notification.GetChannelTypes)
|
||||
notify.GET("/channels", c.Notification.GetChannels)
|
||||
notify.POST("/channels", c.Notification.SaveChannel)
|
||||
notify.DELETE("/channels/:id", c.Notification.DeleteChannel)
|
||||
notify.POST("/channels/test", c.Notification.TestChannel)
|
||||
notify.GET("/bindings", c.Notification.GetBindings)
|
||||
notify.POST("/bindings", c.Notification.SaveBinding)
|
||||
notify.POST("/bindings/batch", c.Notification.BatchSaveBindings)
|
||||
notify.DELETE("/bindings/:id", c.Notification.DeleteBinding)
|
||||
}
|
||||
}
|
||||
|
||||
func registerAppLogRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
appLogs := g.Group("/app-logs")
|
||||
{
|
||||
appLogs.GET("", c.AppLog.GetLogs)
|
||||
appLogs.POST("/read", c.AppLog.MarkAsRead)
|
||||
appLogs.POST("/clear", c.AppLog.ClearLogs)
|
||||
}
|
||||
}
|
||||
|
||||
func registerSystemWSRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
g.GET("/ws/events", c.SystemWS.HandleEvents)
|
||||
}
|
||||
|
||||
func registerMonitorRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
monitor := g.Group("/monitor")
|
||||
{
|
||||
monitor.GET("", c.Monitor.GetSystemMonitor)
|
||||
monitor.GET("/sse", c.Monitor.MonitorSSE)
|
||||
}
|
||||
}
|
||||
|
||||
func initAgentAPIRoutes(root *gin.RouterGroup, c *Controllers) {
|
||||
// Agent API(供远程 Agent 调用,不使用 /v1 版本号)
|
||||
agentAPI := root.Group("/api/agent")
|
||||
{
|
||||
agentAPI.POST("/heartbeat", c.Agent.Heartbeat)
|
||||
agentAPI.GET("/tasks", c.Agent.GetTasks)
|
||||
agentAPI.POST("/report", c.Agent.ReportResult)
|
||||
agentAPI.GET("/download", c.Agent.Download) // 也在这里注册,兼容 Agent 调用
|
||||
agentAPI.GET("/ws", c.Agent.WSConnect) // WebSocket 连接
|
||||
}
|
||||
}
|
||||
|
||||
func registerWebUIRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
webuiGroup := g.Group("/webui")
|
||||
{
|
||||
webuiGroup.GET("", c.WebUI.GetWebUIs)
|
||||
webuiGroup.POST("/upload", c.WebUI.UploadWebUI)
|
||||
webuiGroup.PUT("/active", c.WebUI.SetActiveWebUI)
|
||||
webuiGroup.DELETE("/:name", c.WebUI.DeleteWebUI)
|
||||
}
|
||||
}
|
||||
|
||||
func registerInterconnectRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
interconnect := g.Group("/interconnect")
|
||||
{
|
||||
interconnect.GET("/nodes", c.Interconnect.GetNodes)
|
||||
interconnect.POST("/nodes", c.Interconnect.CreateNode)
|
||||
interconnect.PUT("/nodes/:id", c.Interconnect.UpdateNode)
|
||||
interconnect.DELETE("/nodes/:id", c.Interconnect.DeleteNode)
|
||||
interconnect.GET("/nodes/:id/status", c.Interconnect.GetNodeStatus)
|
||||
interconnect.POST("/sync/script", c.Interconnect.SyncScript)
|
||||
interconnect.POST("/sync/env", c.Interconnect.SyncEnv)
|
||||
interconnect.POST("/sync/task", c.Interconnect.SyncTask)
|
||||
|
||||
interconnect.GET("/child/status", c.Interconnect.GetChildStatus)
|
||||
|
||||
// 代理模式 (面板穿越)
|
||||
interconnect.Any("/proxy/:node_id/*path", c.Interconnect.ProxyRequest)
|
||||
}
|
||||
}
|
||||
|
||||
func registerSystemRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
systemAPI := g.Group("/system")
|
||||
{
|
||||
systemAPI.POST("/export", c.Data.ExportBusinessData)
|
||||
systemAPI.POST("/import", c.Data.ImportBusinessData)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
// "fmt"
|
||||
|
||||
// "github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/eventbus"
|
||||
// "github.com/engigu/taskpool/internal/logger"
|
||||
// "github.com/engigu/taskpool/internal/models"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/executor"
|
||||
)
|
||||
|
||||
func setupEventHandlers(subscribers ...eventbus.Subscriber) {
|
||||
bus := eventbus.DefaultBus
|
||||
|
||||
// 遍历并统一初始化所有订阅者的事件链路
|
||||
for _, s := range subscribers {
|
||||
s.SubscribeEvents(bus)
|
||||
}
|
||||
}
|
||||
|
||||
func startAppLogCleanup(appLogSvc *services.AppLogService) {
|
||||
// 注册到内部系统定时器(并立即执行第一次)
|
||||
executor.GetSysCron().AddJobWithRun("@every 1h", func() {
|
||||
appLogSvc.CleanUp()
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/middleware"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// initOpenAPIV1Routes 初始化 OpenAPI v1 路由
|
||||
// 只注册有 @Tags OpenAPI 注释的接口
|
||||
func initOpenAPIV1Routes(root *gin.RouterGroup, c *Controllers) {
|
||||
// OpenAPI v1 路由组 (使用 Bearer Token)
|
||||
open := root.Group("/open2api/v1")
|
||||
open.Use(middleware.OpenapiRequired())
|
||||
{
|
||||
// 任务相关接口
|
||||
registerOpenAPITaskRoutes(open, c)
|
||||
// 环境变量相关接口
|
||||
registerOpenAPIEnvRoutes(open, c)
|
||||
// 脚本相关接口
|
||||
registerOpenAPIScriptRoutes(open, c)
|
||||
// 日志相关接口
|
||||
registerOpenAPILogRoutes(open, c)
|
||||
// 任务执行相关接口
|
||||
registerOpenAPIExecutorRoutes(open, c)
|
||||
}
|
||||
}
|
||||
|
||||
// registerOpenAPITaskRoutes 注册 OpenAPI 任务路由(只包含有 @Tags OpenAPI 注释的接口)
|
||||
func registerOpenAPITaskRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
tasks := g.Group("/tasks")
|
||||
{
|
||||
tasks.POST("", c.Task.CreateTask)
|
||||
tasks.GET("", c.Task.GetTasks)
|
||||
tasks.GET("/:id", c.Task.GetTask)
|
||||
tasks.PUT("/:id", c.Task.UpdateTask)
|
||||
tasks.DELETE("/:id", c.Task.DeleteTask)
|
||||
tasks.POST("/stop/:logID", c.Task.StopTask)
|
||||
tasks.GET("/tags", c.Task.GetTags)
|
||||
}
|
||||
}
|
||||
|
||||
// registerOpenAPIEnvRoutes 注册 OpenAPI 环境变量路由(只包含有 @Tags OpenAPI 注释的接口)
|
||||
func registerOpenAPIEnvRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
env := g.Group("/env")
|
||||
{
|
||||
env.POST("", c.Env.CreateEnvVar)
|
||||
env.GET("", c.Env.GetEnvVars)
|
||||
env.GET("/all", c.Env.GetAllEnvVars)
|
||||
env.GET("/:id", c.Env.GetEnvVar)
|
||||
env.GET("/:id/tasks", c.Env.GetAssociatedTasks)
|
||||
env.PUT("/:id", c.Env.UpdateEnvVar)
|
||||
env.DELETE("/:id", c.Env.DeleteEnvVar)
|
||||
}
|
||||
}
|
||||
|
||||
// registerOpenAPILogRoutes 注册 OpenAPI 日志路由(只包含有 @Tags OpenAPI 注释的接口)
|
||||
func registerOpenAPILogRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
logs := g.Group("/logs")
|
||||
{
|
||||
logs.GET("", c.Log.GetLogs)
|
||||
logs.GET("/:id", c.Log.GetLogDetail)
|
||||
}
|
||||
}
|
||||
|
||||
// registerOpenAPIExecutorRoutes 注册 OpenAPI 任务执行路由(只包含有 @Tags OpenAPI 注释的接口)
|
||||
func registerOpenAPIExecutorRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
execution := g.Group("/execute")
|
||||
{
|
||||
execution.POST("/task/:id", c.Executor.ExecuteTask)
|
||||
execution.GET("/results", c.Executor.GetLastResults)
|
||||
}
|
||||
}
|
||||
|
||||
// registerOpenAPIScriptRoutes 注册 OpenAPI 脚本路由
|
||||
func registerOpenAPIScriptRoutes(g *gin.RouterGroup, c *Controllers) {
|
||||
scripts := g.Group("/scripts")
|
||||
{
|
||||
scripts.POST("", c.Script.CreateScript)
|
||||
scripts.GET("", c.Script.GetScripts)
|
||||
scripts.GET("/:id", c.Script.GetScript)
|
||||
scripts.PUT("/:id", c.Script.UpdateScript)
|
||||
scripts.DELETE("/:id", c.Script.DeleteScript)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/controllers"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/services/tasks"
|
||||
)
|
||||
|
||||
var executorService *tasks.ExecutorService
|
||||
|
||||
func RegisterControllers() *Controllers {
|
||||
// 初始化服务
|
||||
settingsService := services.NewSettingsService()
|
||||
loginLogService := services.NewLoginLogService()
|
||||
|
||||
// 执行系统初始化(返回 userService)
|
||||
initService := services.NewInitService(settingsService)
|
||||
userService := initService.Initialize()
|
||||
|
||||
taskService := tasks.NewTaskService()
|
||||
envService := services.NewEnvService()
|
||||
scriptService := services.NewScriptService()
|
||||
sendStatsService := services.NewSendStatsService()
|
||||
agentWSManager := services.GetAgentWSManager()
|
||||
systemWSManager := services.GetSystemWSManager()
|
||||
|
||||
taskLogService := tasks.NewTaskLogService(sendStatsService)
|
||||
// 创建任务执行服务(需要依赖注入)
|
||||
notifyService := services.NewNotificationService()
|
||||
appLogService := services.NewAppLogService()
|
||||
interconnectService := services.NewInterconnectService()
|
||||
|
||||
// 清理 task 运行状态的任务可以直接由 executorService 承担或在此处通过 Database 直接清理
|
||||
// 简单期间,我们使用一个新方法 tasks.CleanupRunningTasks() 或者让 executorService 启动时清理
|
||||
|
||||
executorService = tasks.NewExecutorService(taskService, taskLogService, agentWSManager, settingsService, envService)
|
||||
// 启动时清理残留的运行状态
|
||||
_ = executorService.CleanupRunningTasks()
|
||||
|
||||
// 启动计划任务
|
||||
executorService.StartCron()
|
||||
|
||||
// 初始化所有关注系统总线的服务
|
||||
setupEventHandlers(appLogService, notifyService, loginLogService, systemWSManager)
|
||||
startAppLogCleanup(appLogService)
|
||||
|
||||
taskController := controllers.NewTaskController(taskService, executorService)
|
||||
envController := controllers.NewEnvController(envService)
|
||||
|
||||
// 初始化并返回控制器
|
||||
return &Controllers{
|
||||
Task: taskController,
|
||||
Auth: controllers.NewAuthController(userService, settingsService, loginLogService),
|
||||
Env: envController,
|
||||
Script: controllers.NewScriptController(scriptService),
|
||||
Executor: controllers.NewExecutorController(executorService),
|
||||
File: controllers.NewFileController(constant.ScriptsWorkDir),
|
||||
Dashboard: controllers.NewDashboardController(executorService),
|
||||
Log: controllers.NewLogController(),
|
||||
LogSSE: controllers.NewLogSSEController(),
|
||||
Terminal: controllers.NewTerminalController(envService),
|
||||
Settings: controllers.NewSettingsController(userService, loginLogService, executorService),
|
||||
Dependency: controllers.NewDependencyController(),
|
||||
Agent: controllers.NewAgentController(settingsService),
|
||||
Mise: controllers.NewMiseController(services.NewMiseService()),
|
||||
Notification: controllers.NewNotificationController(),
|
||||
AppLog: controllers.NewAppLogController(),
|
||||
SystemWS: controllers.NewSystemWSController(),
|
||||
WebUI: controllers.NewWebUIController(services.NewWebUIService(settingsService)),
|
||||
Monitor: controllers.NewMonitorController(executorService),
|
||||
Interconnect: controllers.NewInterconnectController(interconnectService),
|
||||
Data: controllers.NewDataController(taskController, envController),
|
||||
Install: controllers.NewInstallController(),
|
||||
}
|
||||
}
|
||||
|
||||
// StopCron 停止计划任务服务
|
||||
func StopCron() {
|
||||
if executorService != nil {
|
||||
executorService.Stop()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/controllers"
|
||||
"github.com/engigu/taskpool/internal/middleware"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
|
||||
"github.com/gin-contrib/pprof"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type Controllers struct {
|
||||
Task *controllers.TaskController
|
||||
Auth *controllers.AuthController
|
||||
Env *controllers.EnvController
|
||||
Script *controllers.ScriptController
|
||||
Executor *controllers.ExecutorController
|
||||
File *controllers.FileController
|
||||
Dashboard *controllers.DashboardController
|
||||
Log *controllers.LogController
|
||||
LogSSE *controllers.LogSSEController
|
||||
Terminal *controllers.TerminalController
|
||||
Settings *controllers.SettingsController
|
||||
Dependency *controllers.DependencyController
|
||||
Agent *controllers.AgentController
|
||||
Mise *controllers.MiseController
|
||||
Notification *controllers.NotificationController
|
||||
AppLog *controllers.AppLogController
|
||||
SystemWS *controllers.SystemWSController
|
||||
WebUI *controllers.WebUIController
|
||||
Monitor *controllers.MonitorController
|
||||
Interconnect *controllers.InterconnectController
|
||||
Data *controllers.DataController
|
||||
Install *controllers.InstallController
|
||||
}
|
||||
|
||||
func Setup(c *Controllers) *gin.Engine {
|
||||
if os.Getenv("GIN_MODE") == "" {
|
||||
gin.SetMode(gin.ReleaseMode)
|
||||
}
|
||||
router := gin.New()
|
||||
router.Use(middleware.GinLogger(), middleware.GinRecovery())
|
||||
router.Use(middleware.TravelProxyMiddleware())
|
||||
|
||||
// 获取 URL 前缀
|
||||
cfg := services.GetConfig()
|
||||
urlPrefix := strings.TrimSuffix(cfg.Server.URLPrefix, "/")
|
||||
|
||||
// 创建一个路由组,如果有前缀则使用前缀,否则使用根路径
|
||||
var root *gin.RouterGroup
|
||||
if urlPrefix != "" {
|
||||
root = router.Group(urlPrefix)
|
||||
} else {
|
||||
root = router.Group("")
|
||||
}
|
||||
|
||||
// 按需绑定 Pprof 调试路由 (注册在 root 下以支持 URLPrefix)
|
||||
if cfg.Server.PprofEnabled {
|
||||
// pprof.RouteRegister 会在传入的路由组下注册 /debug/pprof 等路由
|
||||
pprof.RouteRegister(root)
|
||||
}
|
||||
|
||||
// =========================================================================
|
||||
// 路由分类组装 (对应 Nginx 的 location 块分发)
|
||||
// =========================================================================
|
||||
|
||||
// 1. [ location /assets ] 静态资源路由
|
||||
initStaticRoutes(root)
|
||||
|
||||
// 3. [ location /api ] 内部 API 路由组
|
||||
apiV1 := root.Group("/api/v1")
|
||||
initPublicAPIRoutes(apiV1, c) // 公开接口 (无需认证)
|
||||
initAuthorizedAPIRoutes(apiV1, c) // 授权接口 (需 JWT)
|
||||
|
||||
// 4. [ location /api/agent ] Agent 相关 API 路由组
|
||||
initAgentAPIRoutes(root, c)
|
||||
initOpenAPIV1Routes(root, c)
|
||||
|
||||
// =========================================================================
|
||||
// [ location / ] 全局 404 兜底与 SPA 渲染
|
||||
// 对应 Nginx: try_files $uri $uri/ /index.html;
|
||||
// =========================================================================
|
||||
router.NoRoute(func(ctx *gin.Context) {
|
||||
path := ctx.Request.URL.Path
|
||||
|
||||
// 如果配置了前缀,只处理带前缀的路径
|
||||
if urlPrefix != "" && !strings.HasPrefix(path, urlPrefix) {
|
||||
ctx.Status(404)
|
||||
return
|
||||
}
|
||||
|
||||
// 解析实际的相对路径
|
||||
relPath := strings.TrimPrefix(path, urlPrefix)
|
||||
if !strings.HasPrefix(relPath, "/") {
|
||||
relPath = "/" + relPath
|
||||
}
|
||||
|
||||
// 拦截器:不该返回 index.html 的情况
|
||||
// 如果该请求被识别为 API 请求、静态资源请求,或者是带有明确文件后缀(如 .js / .css / .png)的物理文件请求
|
||||
// 都不应该返回 SPA 页面(会报前端 MIME 类型错误),而是直接掐断返回 404
|
||||
hasAnyExt := false
|
||||
if idx := strings.LastIndex(relPath, "."); idx > 0 && len(relPath)-idx < 6 {
|
||||
// 简单判断是否有后缀(如 .js, .css)
|
||||
hasAnyExt = true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(relPath, "/api/") || strings.HasPrefix(relPath, "/assets/") || strings.HasPrefix(relPath, "/debug/") || hasAnyExt {
|
||||
ctx.String(404, "404 Not Found")
|
||||
return
|
||||
}
|
||||
|
||||
// 其他所有有效的前端页面路径(如 /tasks, /settings),都返回 index.html 交给 vue-router 处理
|
||||
serveSPA(ctx, urlPrefix, 200)
|
||||
})
|
||||
|
||||
return router
|
||||
}
|
||||
@@ -0,0 +1,268 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"compress/gzip"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"io/fs"
|
||||
"mime"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
"github.com/engigu/taskpool/internal/services"
|
||||
"github.com/engigu/taskpool/internal/static"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// cacheControl 返回设置 Cache-Control header 的中间件
|
||||
func cacheControl(value string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.Header("Cache-Control", value)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func openFileWithWebui(filename string) (fs.File, error) {
|
||||
webuiSvc := services.NewWebUIService(services.NewSettingsService())
|
||||
if customFS := webuiSvc.GetActiveWebUIFS(); customFS != nil {
|
||||
// 如果启用了定义的前端包,去取定义的路径
|
||||
return customFS.Open(filename)
|
||||
}
|
||||
|
||||
// 如果是默认的,取默认路径
|
||||
defaultFS := static.GetFS()
|
||||
if defaultFS == nil {
|
||||
return nil, fs.ErrNotExist
|
||||
}
|
||||
return defaultFS.Open(filename)
|
||||
}
|
||||
|
||||
func readFileWithWebui(filename string) ([]byte, error) {
|
||||
webuiSvc := services.NewWebUIService(services.NewSettingsService())
|
||||
if customFS := webuiSvc.GetActiveWebUIFS(); customFS != nil {
|
||||
// 如果启用了定义的前端包,去取定义的路径
|
||||
return fs.ReadFile(customFS, filename)
|
||||
}
|
||||
|
||||
// 如果是默认的,取默认路径
|
||||
defaultFS := static.GetFS()
|
||||
if defaultFS == nil {
|
||||
return nil, fs.ErrNotExist
|
||||
}
|
||||
return fs.ReadFile(defaultFS, filename)
|
||||
}
|
||||
|
||||
func initStaticRoutes(root *gin.RouterGroup) {
|
||||
|
||||
// 专门处理 /assets 目录下的资源
|
||||
root.GET("/assets/*filepath", cacheControl("public, max-age=31536000, immutable"), func(ctx *gin.Context) {
|
||||
fullPath := "assets" + ctx.Param("filepath")
|
||||
fullPath = strings.TrimPrefix(fullPath, "/")
|
||||
|
||||
isGzipSupported := strings.Contains(ctx.GetHeader("Accept-Encoding"), "gzip")
|
||||
gzPath := fullPath + ".gz"
|
||||
|
||||
// 确定 MIME 类型
|
||||
ext := filepath.Ext(fullPath)
|
||||
contentType := mime.TypeByExtension(ext)
|
||||
if contentType == "" {
|
||||
switch ext {
|
||||
case ".js":
|
||||
contentType = "application/javascript"
|
||||
case ".css":
|
||||
contentType = "text/css"
|
||||
case ".svg":
|
||||
contentType = "image/svg+xml"
|
||||
default:
|
||||
contentType = "application/octet-stream"
|
||||
}
|
||||
}
|
||||
|
||||
// 优先尝试读取 .gz 文件
|
||||
if gzFile, err := openFileWithWebui(gzPath); err == nil {
|
||||
defer gzFile.Close()
|
||||
ctx.Header("Content-Type", contentType)
|
||||
|
||||
if isGzipSupported {
|
||||
// 极致性能:流式透传压缩包 (RSS 占用极低)
|
||||
ctx.Header("Content-Encoding", "gzip")
|
||||
ctx.Status(http.StatusOK)
|
||||
io.Copy(ctx.Writer, gzFile)
|
||||
} else {
|
||||
// 兼容处理:流式解压发送
|
||||
gr, _ := gzip.NewReader(gzFile)
|
||||
defer gr.Close()
|
||||
ctx.Status(http.StatusOK)
|
||||
io.Copy(ctx.Writer, gr)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// 如果没有 .gz,流式读取原文件
|
||||
if file, err := openFileWithWebui(fullPath); err == nil {
|
||||
defer file.Close()
|
||||
ctx.Header("Content-Type", contentType)
|
||||
ctx.Status(http.StatusOK)
|
||||
io.Copy(ctx.Writer, file)
|
||||
return
|
||||
}
|
||||
|
||||
ctx.Status(404)
|
||||
})
|
||||
|
||||
// logo.svg 等单文件处理
|
||||
root.GET("/logo.svg", func(ctx *gin.Context) {
|
||||
settings := services.NewSettingsService()
|
||||
icon := settings.Get(constant.SectionSite, constant.KeyIcon)
|
||||
if icon != "" {
|
||||
ctx.Header("Cache-Control", "public, max-age=86400")
|
||||
ctx.Data(http.StatusOK, "image/svg+xml", []byte(icon))
|
||||
return
|
||||
}
|
||||
serveSingleFile(ctx, "logo.svg", "image/svg+xml", "public, max-age=86400")
|
||||
})
|
||||
|
||||
// PWA 相关路由处理
|
||||
initPWARoutes(root)
|
||||
}
|
||||
|
||||
func initPWARoutes(root *gin.RouterGroup) {
|
||||
// PWA 相关文件处理
|
||||
pwaRootFiles := map[string]string{
|
||||
"/sw.js": "application/javascript",
|
||||
"/registerSW.js": "application/javascript",
|
||||
"/favicon.ico": "image/x-icon",
|
||||
"/pwa-icon-192.png": "image/png",
|
||||
"/pwa-icon-512.png": "image/png",
|
||||
}
|
||||
|
||||
for path, contentType := range pwaRootFiles {
|
||||
pPath := path
|
||||
pType := contentType
|
||||
root.GET(pPath, func(ctx *gin.Context) {
|
||||
file := strings.TrimPrefix(pPath, "/")
|
||||
serveSingleFile(ctx, file, pType, "public, no-cache")
|
||||
})
|
||||
}
|
||||
|
||||
// 动态 manifest 处理 (支持由 Go 后端控制标题和图标)
|
||||
root.GET("/manifest.webmanifest", handleManifest)
|
||||
|
||||
// 动态匹配 workbox-*.js (Vite PWA 生成的库文件)
|
||||
root.GET("/workbox-:hash.js", func(ctx *gin.Context) {
|
||||
file := "workbox-" + ctx.Param("hash") + ".js"
|
||||
serveSingleFile(ctx, file, "application/javascript", "public, max-age=31536000, immutable")
|
||||
})
|
||||
}
|
||||
|
||||
func handleManifest(ctx *gin.Context) {
|
||||
|
||||
// 读取原始 manifest
|
||||
data, err := readFileWithWebui("manifest.webmanifest")
|
||||
if err != nil {
|
||||
ctx.Status(404)
|
||||
return
|
||||
}
|
||||
|
||||
var manifest map[string]interface{}
|
||||
if err := json.Unmarshal(data, &manifest); err != nil {
|
||||
// 如果解析失败,回退到原始文件
|
||||
ctx.Data(200, "application/manifest+json", data)
|
||||
return
|
||||
}
|
||||
|
||||
// 注入后端配置的标题
|
||||
settings := services.NewSettingsService()
|
||||
title := settings.Get(constant.SectionSite, constant.KeyTitle)
|
||||
if title != "" {
|
||||
manifest["name"] = title
|
||||
manifest["short_name"] = title
|
||||
}
|
||||
|
||||
// 注入后端配置的图标 (首选 logo.svg)
|
||||
manifest["icons"] = []map[string]interface{}{
|
||||
{
|
||||
"src": "/logo.svg",
|
||||
"sizes": "any",
|
||||
"type": "image/svg+xml",
|
||||
"purpose": "any maskable",
|
||||
},
|
||||
}
|
||||
|
||||
res, _ := json.Marshal(manifest)
|
||||
ctx.Header("Cache-Control", "public, no-cache")
|
||||
ctx.Data(200, "application/manifest+json", res)
|
||||
}
|
||||
|
||||
func serveSingleFile(ctx *gin.Context, filename string, contentType string, cache string) {
|
||||
|
||||
if cache != "" {
|
||||
ctx.Header("Cache-Control", cache)
|
||||
}
|
||||
ctx.Header("Content-Type", contentType)
|
||||
|
||||
isGzipSupported := strings.Contains(ctx.GetHeader("Accept-Encoding"), "gzip")
|
||||
|
||||
// 尝试流式发送压缩版
|
||||
if gzFile, err := openFileWithWebui(filename + ".gz"); err == nil {
|
||||
defer gzFile.Close()
|
||||
if isGzipSupported {
|
||||
ctx.Header("Content-Encoding", "gzip")
|
||||
ctx.Status(200)
|
||||
io.Copy(ctx.Writer, gzFile)
|
||||
} else {
|
||||
gr, _ := gzip.NewReader(gzFile)
|
||||
defer gr.Close()
|
||||
ctx.Status(200)
|
||||
io.Copy(ctx.Writer, gr)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// 尝试流式发送原版
|
||||
if file, err := openFileWithWebui(filename); err == nil {
|
||||
defer file.Close()
|
||||
ctx.Status(200)
|
||||
io.Copy(ctx.Writer, file)
|
||||
return
|
||||
}
|
||||
|
||||
ctx.Status(404)
|
||||
}
|
||||
|
||||
// serveSPA 注入配置并返回 index.html 给前端渲染
|
||||
func serveSPA(ctx *gin.Context, urlPrefix string, status int) {
|
||||
|
||||
var data []byte
|
||||
// index.html 较小且需要修改字符串,可以一次性读入内存
|
||||
if gzFile, err := openFileWithWebui("index.html.gz"); err == nil {
|
||||
defer gzFile.Close()
|
||||
gr, _ := gzip.NewReader(gzFile)
|
||||
data, _ = io.ReadAll(gr)
|
||||
gr.Close()
|
||||
} else if file, err := openFileWithWebui("index.html"); err == nil {
|
||||
defer file.Close()
|
||||
data, _ = io.ReadAll(file)
|
||||
}
|
||||
|
||||
if data == nil {
|
||||
ctx.String(status, "index.html not found.")
|
||||
return
|
||||
}
|
||||
|
||||
html := string(data)
|
||||
baseHref := urlPrefix + "/"
|
||||
if urlPrefix == "" {
|
||||
baseHref = "/"
|
||||
}
|
||||
html = strings.Replace(html, "<head>", "<head>\n <base href=\""+baseHref+"\">", 1)
|
||||
configScript := `<script>window.__BASE_URL__ = "` + urlPrefix + `"; window.__API_VERSION__ = "/api/v1";</script>`
|
||||
html = strings.Replace(html, "</head>", configScript+"</head>", 1)
|
||||
|
||||
ctx.Header("Content-Type", "text/html; charset=utf-8")
|
||||
ctx.Header("Cache-Control", "no-cache, no-store, must-revalidate")
|
||||
ctx.Data(status, "text/html; charset=utf-8", []byte(html))
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/net/proxy"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type barkResponse struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
type Bark struct {
|
||||
PushKey string
|
||||
Archive string
|
||||
Group string
|
||||
Sound string
|
||||
Icon string
|
||||
Level string
|
||||
URL string
|
||||
Key string
|
||||
IV string
|
||||
Server string
|
||||
Badge string
|
||||
Copy string
|
||||
AutoCopy string
|
||||
ProxyURL string // 可选的代理地址
|
||||
}
|
||||
|
||||
func (b *Bark) Request(title, content string) ([]byte, error) {
|
||||
data := map[string]interface{}{
|
||||
"device_key": b.PushKey,
|
||||
"title": title,
|
||||
"body": content,
|
||||
}
|
||||
if b.Archive != "" {
|
||||
data["isArchive"] = b.Archive
|
||||
}
|
||||
if b.Group != "" {
|
||||
data["group"] = b.Group
|
||||
}
|
||||
if b.Sound != "" {
|
||||
data["sound"] = b.Sound
|
||||
}
|
||||
if b.Icon != "" {
|
||||
data["icon"] = b.Icon
|
||||
}
|
||||
if b.Level != "" {
|
||||
data["level"] = b.Level
|
||||
}
|
||||
if b.URL != "" {
|
||||
data["url"] = b.URL
|
||||
}
|
||||
if b.Badge != "" {
|
||||
data["badge"] = b.Badge
|
||||
}
|
||||
if b.Copy != "" {
|
||||
data["copy"] = b.Copy
|
||||
}
|
||||
if b.AutoCopy != "" {
|
||||
data["autoCopy"] = b.AutoCopy
|
||||
}
|
||||
|
||||
server := b.Server
|
||||
if server == "" {
|
||||
server = "https://api.day.app"
|
||||
}
|
||||
server = strings.TrimSuffix(server, "/")
|
||||
apiURL := server + "/push"
|
||||
|
||||
// If PushKey is a full URL, we might be using an old-style custom URL
|
||||
if strings.HasPrefix(b.PushKey, "http") {
|
||||
apiURL = b.PushKey
|
||||
}
|
||||
|
||||
var postData interface{}
|
||||
if b.Key != "" && b.IV != "" {
|
||||
// Encrypted Request
|
||||
// 1. Prepare the full notification payload (without device_key, as specified for encryption)
|
||||
encryptData := make(map[string]interface{})
|
||||
for k, v := range data {
|
||||
if k != "device_key" {
|
||||
encryptData[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(encryptData)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ciphertext, err := b.encryptPayload(string(jsonData))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encryption failed: %v", err)
|
||||
}
|
||||
|
||||
postData = map[string]interface{}{
|
||||
"ciphertext": ciphertext,
|
||||
"iv": b.IV,
|
||||
"device_key": b.PushKey,
|
||||
}
|
||||
} else {
|
||||
// Normal request
|
||||
postData = data
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(postData)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 使用带超时的客户端
|
||||
client := b.getHTTPClient()
|
||||
resp, err := client.Post(apiURL, "application/json;charset=utf-8", bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var r barkResponse
|
||||
err = json.Unmarshal(body, &r)
|
||||
if err != nil {
|
||||
// If not JSON, return the raw body as it might be a simple success message from some servers
|
||||
if resp.StatusCode == 200 {
|
||||
return body, nil
|
||||
}
|
||||
return body, err
|
||||
}
|
||||
|
||||
if r.Code != 200 && resp.StatusCode != 200 {
|
||||
return body, fmt.Errorf("bark response error: %s", string(body))
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
func (b *Bark) encryptPayload(payload string) (string, error) {
|
||||
key := []byte(b.Key)
|
||||
iv := []byte(b.IV)
|
||||
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
paddedPayload := b.pkcs7Pad([]byte(payload), aes.BlockSize)
|
||||
mode := cipher.NewCBCEncrypter(block, iv)
|
||||
ciphertext := make([]byte, len(paddedPayload))
|
||||
mode.CryptBlocks(ciphertext, paddedPayload)
|
||||
|
||||
return base64.StdEncoding.EncodeToString(ciphertext), nil
|
||||
}
|
||||
|
||||
func (b *Bark) pkcs7Pad(data []byte, blockSize int) []byte {
|
||||
padding := blockSize - len(data)%blockSize
|
||||
padtext := bytes.Repeat([]byte{byte(padding)}, padding)
|
||||
return append(data, padtext...)
|
||||
}
|
||||
|
||||
// getHTTPClient 获取 HTTP 客户端(含超时和代理)
|
||||
func (b *Bark) getHTTPClient() *http.Client {
|
||||
client := &http.Client{
|
||||
Timeout: 20 * time.Second,
|
||||
}
|
||||
|
||||
if b.ProxyURL != "" {
|
||||
proxyURL, err := url.Parse(b.ProxyURL)
|
||||
if err == nil {
|
||||
if strings.HasPrefix(strings.ToLower(b.ProxyURL), "socks5://") {
|
||||
dialer, err := b.createSOCKS5Dialer(proxyURL)
|
||||
if err == nil {
|
||||
client.Transport = &http.Transport{
|
||||
DialContext: dialer.DialContext,
|
||||
}
|
||||
}
|
||||
} else {
|
||||
client.Transport = &http.Transport{
|
||||
Proxy: http.ProxyURL(proxyURL),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return client
|
||||
}
|
||||
|
||||
// createSOCKS5Dialer 创建 SOCKS5 拨号器
|
||||
func (b *Bark) createSOCKS5Dialer(proxyURL *url.URL) (proxy.ContextDialer, error) {
|
||||
host := proxyURL.Host
|
||||
var auth *proxy.Auth
|
||||
if proxyURL.User != nil {
|
||||
password, _ := proxyURL.User.Password()
|
||||
auth = &proxy.Auth{
|
||||
User: proxyURL.User.Username(),
|
||||
Password: password,
|
||||
}
|
||||
}
|
||||
|
||||
baseDialer := &net.Dialer{
|
||||
Timeout: 20 * time.Second,
|
||||
KeepAlive: 20 * time.Second,
|
||||
}
|
||||
|
||||
dialer, err := proxy.SOCKS5("tcp", host, auth, baseDialer)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
contextDialer, ok := dialer.(proxy.ContextDialer)
|
||||
if !ok {
|
||||
return nil, errors.New("failed to convert to ContextDialer")
|
||||
}
|
||||
|
||||
return contextDialer, nil
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type CustomWebhook struct {
|
||||
Webhook string
|
||||
Body string
|
||||
}
|
||||
|
||||
var Client = &http.Client{
|
||||
Timeout: 5 * time.Second,
|
||||
}
|
||||
|
||||
func (cw *CustomWebhook) Request(url string, msg string, headers map[string]string) ([]byte, error) {
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer([]byte(msg)))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
|
||||
resp, err := Client.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func(Body io.ReadCloser) {
|
||||
err := Body.Close()
|
||||
if err != nil {
|
||||
|
||||
}
|
||||
}(resp.Body)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return body, err
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type response struct {
|
||||
Code int `json:"errcode"`
|
||||
Msg string `json:"errmsg"`
|
||||
}
|
||||
|
||||
type Dtalk struct {
|
||||
AccessToken string
|
||||
Secret string
|
||||
}
|
||||
|
||||
func (t *Dtalk) Request(msg interface{}) ([]byte, error) {
|
||||
b, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := http.Post(t.getURL(), "application/json", bytes.NewBuffer(b))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
defer func(Body io.ReadCloser) {
|
||||
err := Body.Close()
|
||||
if err != nil {
|
||||
|
||||
}
|
||||
}(resp.Body)
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var r response
|
||||
err = json.Unmarshal(body, &r)
|
||||
if err != nil {
|
||||
return body, err
|
||||
}
|
||||
if r.Code != 0 {
|
||||
return body, fmt.Errorf("response error: %s", string(body))
|
||||
}
|
||||
return body, err
|
||||
}
|
||||
|
||||
// SendMessageText Function to send message
|
||||
func (t *Dtalk) SendMessageText(text string, at ...string) ([]byte, error) {
|
||||
msg := map[string]interface{}{
|
||||
"msgtype": "text",
|
||||
"text": map[string]string{
|
||||
"content": text,
|
||||
},
|
||||
}
|
||||
|
||||
// 添加@功能
|
||||
if len(at) > 0 {
|
||||
atMobiles := []string{}
|
||||
isAtAll := false
|
||||
|
||||
for _, mobile := range at {
|
||||
if mobile == "all" || mobile == "@all" {
|
||||
isAtAll = true
|
||||
} else {
|
||||
atMobiles = append(atMobiles, mobile)
|
||||
}
|
||||
}
|
||||
|
||||
msg["at"] = map[string]interface{}{
|
||||
"atMobiles": atMobiles,
|
||||
"isAtAll": isAtAll,
|
||||
}
|
||||
}
|
||||
|
||||
resp, err := t.Request(msg)
|
||||
return resp, err
|
||||
}
|
||||
|
||||
func (t *Dtalk) SendMessageMarkdown(title, text string, at ...string) ([]byte, error) {
|
||||
msg := map[string]interface{}{
|
||||
"msgtype": "markdown",
|
||||
"markdown": map[string]string{
|
||||
"title": title,
|
||||
"text": text,
|
||||
},
|
||||
}
|
||||
|
||||
// 添加@功能
|
||||
if len(at) > 0 {
|
||||
atMobiles := []string{}
|
||||
isAtAll := false
|
||||
|
||||
for _, mobile := range at {
|
||||
if mobile == "all" || mobile == "@all" {
|
||||
isAtAll = true
|
||||
} else {
|
||||
atMobiles = append(atMobiles, mobile)
|
||||
}
|
||||
}
|
||||
|
||||
msg["at"] = map[string]interface{}{
|
||||
"atMobiles": atMobiles,
|
||||
"isAtAll": isAtAll,
|
||||
}
|
||||
}
|
||||
|
||||
resp, err := t.Request(msg)
|
||||
return resp, err
|
||||
}
|
||||
|
||||
func (t *Dtalk) hmacSha256(stringToSign string, secret string) string {
|
||||
h := hmac.New(sha256.New, []byte(secret))
|
||||
h.Write([]byte(stringToSign))
|
||||
return base64.StdEncoding.EncodeToString(h.Sum(nil))
|
||||
}
|
||||
|
||||
func (t *Dtalk) getURL() string {
|
||||
wh := "https://oapi.dingtalk.com/robot/send?access_token=" + t.AccessToken
|
||||
timestamp := time.Now().UnixNano() / 1e6
|
||||
stringToSign := fmt.Sprintf("%d\n%s", timestamp, t.Secret)
|
||||
sign := t.hmacSha256(stringToSign, t.Secret)
|
||||
url := fmt.Sprintf("%s×tamp=%d&sign=%s", wh, timestamp, sign)
|
||||
return url
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"gopkg.in/gomail.v2"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type EmailMessage struct {
|
||||
Server string
|
||||
Port int
|
||||
Account string
|
||||
Passwd string
|
||||
FromName string
|
||||
GM *gomail.Dialer
|
||||
}
|
||||
|
||||
func (e *EmailMessage) Init(host string, port int, account string, passwd string, fromName string) {
|
||||
e.Server = host
|
||||
e.Port = port
|
||||
e.Account = account
|
||||
e.Passwd = passwd
|
||||
e.FromName = fromName
|
||||
e.GM = gomail.NewDialer(host, port, account, passwd)
|
||||
}
|
||||
|
||||
func (e *EmailMessage) sendMessage(toEmail string, title string, content string, contentType string) string {
|
||||
m := gomail.NewMessage()
|
||||
if e.FromName != "" {
|
||||
m.SetAddressHeader("From", e.Account, e.FromName)
|
||||
} else {
|
||||
m.SetHeader("From", e.Account)
|
||||
}
|
||||
m.SetHeader("To", toEmail)
|
||||
m.SetHeader("Subject", title)
|
||||
m.SetBody(contentType, content)
|
||||
|
||||
if err := e.GM.DialAndSend(m); err != nil {
|
||||
return fmt.Sprintf("邮件发送失败: %s", err)
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (e *EmailMessage) SendTextMessage(toEmail string, title string, content string) string {
|
||||
return e.sendMessage(toEmail, title, content, "text/plain")
|
||||
}
|
||||
|
||||
func (e *EmailMessage) SendHtmlMessage(toEmail string, title string, content string) string {
|
||||
return e.sendMessage(toEmail, title, content, "text/html")
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type feishuResponse struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
|
||||
type Feishu struct {
|
||||
AccessToken string
|
||||
Secret string
|
||||
}
|
||||
|
||||
// genSign 生成飞书签名
|
||||
func (f *Feishu) genSign(timestamp int64) string {
|
||||
if f.Secret == "" {
|
||||
return ""
|
||||
}
|
||||
stringToSign := fmt.Sprintf("%v\n%s", timestamp, f.Secret)
|
||||
h := hmac.New(sha256.New, []byte(stringToSign))
|
||||
signature := base64.StdEncoding.EncodeToString(h.Sum(nil))
|
||||
return signature
|
||||
}
|
||||
|
||||
// SendMessageText 发送文本消息
|
||||
func (f *Feishu) SendMessageText(content string, atMobiles ...string) ([]byte, error) {
|
||||
timestamp := time.Now().Unix()
|
||||
sign := f.genSign(timestamp)
|
||||
|
||||
msg := map[string]interface{}{
|
||||
"timestamp": strconv.FormatInt(timestamp, 10),
|
||||
"sign": sign,
|
||||
"msg_type": "text",
|
||||
"content": map[string]interface{}{
|
||||
"text": content,
|
||||
},
|
||||
}
|
||||
|
||||
return f.send(msg)
|
||||
}
|
||||
|
||||
// SendMessageMarkdown 发送 Markdown 消息
|
||||
func (f *Feishu) SendMessageMarkdown(title, content string, atMobiles ...string) ([]byte, error) {
|
||||
timestamp := time.Now().Unix()
|
||||
sign := f.genSign(timestamp)
|
||||
|
||||
// 处理 @ 人员
|
||||
atContent := ""
|
||||
if len(atMobiles) > 0 {
|
||||
for _, mobile := range atMobiles {
|
||||
if mobile == "all" {
|
||||
atContent += "<at user_id=\"all\">所有人</at>"
|
||||
} else {
|
||||
atContent += fmt.Sprintf("<at user_id=\"%s\"></at>", mobile)
|
||||
}
|
||||
}
|
||||
content = atContent + "\n" + content
|
||||
}
|
||||
|
||||
msg := map[string]interface{}{
|
||||
"timestamp": strconv.FormatInt(timestamp, 10),
|
||||
"sign": sign,
|
||||
"msg_type": "interactive",
|
||||
"card": map[string]interface{}{
|
||||
"header": map[string]interface{}{
|
||||
"title": map[string]interface{}{
|
||||
"tag": "plain_text",
|
||||
"content": title,
|
||||
},
|
||||
},
|
||||
"elements": []map[string]interface{}{
|
||||
{
|
||||
"tag": "markdown",
|
||||
"content": content,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
return f.send(msg)
|
||||
}
|
||||
|
||||
// send 发送请求
|
||||
func (f *Feishu) send(msg map[string]interface{}) ([]byte, error) {
|
||||
url := fmt.Sprintf("https://open.feishu.cn/open-apis/bot/v2/hook/%s", f.AccessToken)
|
||||
|
||||
jsonData, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("JSON序列化失败: %v", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建请求失败: %v", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := &http.Client{Timeout: 10 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("发送请求失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取响应失败: %v", err)
|
||||
}
|
||||
|
||||
var result feishuResponse
|
||||
if err := json.Unmarshal(body, &result); err != nil {
|
||||
return body, fmt.Errorf("解析响应失败: %v", err)
|
||||
}
|
||||
|
||||
if result.Code != 0 {
|
||||
return body, fmt.Errorf("飞书返回错误: %s", result.Msg)
|
||||
}
|
||||
|
||||
return body, nil
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type gotifyResponse struct {
|
||||
Id int `json:"id"`
|
||||
Message string `json:"message"`
|
||||
ErrorCode int `json:"errorCode"`
|
||||
}
|
||||
|
||||
type Gotify struct {
|
||||
Url string
|
||||
Token string
|
||||
Priority int
|
||||
}
|
||||
|
||||
func (g *Gotify) Request(title, content string) ([]byte, error) {
|
||||
// Construct the URL with token
|
||||
u, err := url.Parse(fmt.Sprintf("%s/message", g.Url))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
q := u.Query()
|
||||
q.Set("token", g.Token)
|
||||
u.RawQuery = q.Encode()
|
||||
|
||||
data := map[string]interface{}{
|
||||
"title": title,
|
||||
"message": content,
|
||||
"priority": g.Priority,
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resp, err := http.Post(u.String(), "application/json", bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var r gotifyResponse
|
||||
err = json.Unmarshal(body, &r)
|
||||
if err != nil {
|
||||
return body, err
|
||||
}
|
||||
|
||||
if r.Id == 0 {
|
||||
return body, fmt.Errorf("gotify response error: %s", string(body))
|
||||
}
|
||||
|
||||
return body, nil
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type Ntfy struct {
|
||||
Url string
|
||||
Topic string
|
||||
Priority string
|
||||
Icon string
|
||||
Token string
|
||||
Username string
|
||||
Password string
|
||||
Actions string
|
||||
}
|
||||
|
||||
func encodeRFC2047(text string) string {
|
||||
encoded := base64.StdEncoding.EncodeToString([]byte(text))
|
||||
return fmt.Sprintf("=?utf-8?B?%s?=", encoded)
|
||||
}
|
||||
|
||||
func (n *Ntfy) Request(title, content string) ([]byte, error) {
|
||||
if n.Url == "" {
|
||||
n.Url = "https://ntfy.sh"
|
||||
}
|
||||
url := fmt.Sprintf("%s/%s", n.Url, n.Topic)
|
||||
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBufferString(content))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
req.Header.Set("Title", encodeRFC2047(title))
|
||||
priority := n.Priority
|
||||
if priority == "" {
|
||||
priority = "3"
|
||||
}
|
||||
req.Header.Set("Priority", priority)
|
||||
if n.Icon != "" {
|
||||
req.Header.Set("Icon", n.Icon)
|
||||
}
|
||||
if n.Actions != "" {
|
||||
req.Header.Set("Actions", encodeRFC2047(n.Actions))
|
||||
}
|
||||
|
||||
if n.Token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+n.Token)
|
||||
} else if n.Username != "" && n.Password != "" {
|
||||
authStr := n.Username + ":" + n.Password
|
||||
encodedAuth := base64.StdEncoding.EncodeToString([]byte(authStr))
|
||||
req.Header.Set("Authorization", "Basic "+encodedAuth)
|
||||
}
|
||||
|
||||
client := &http.Client{}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return body, fmt.Errorf("ntfy response error: %s", string(body))
|
||||
}
|
||||
|
||||
return body, nil
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type PushMe struct {
|
||||
PushKey string
|
||||
URL string
|
||||
Date string
|
||||
Type string
|
||||
}
|
||||
|
||||
func (p *PushMe) Request(title, content string) (string, error) {
|
||||
apiURL := p.URL
|
||||
if apiURL == "" {
|
||||
apiURL = "https://push.i-i.me/"
|
||||
}
|
||||
|
||||
data := url.Values{}
|
||||
data.Set("push_key", p.PushKey)
|
||||
data.Set("title", title)
|
||||
data.Set("content", content)
|
||||
if p.Date != "" {
|
||||
data.Set("date", p.Date)
|
||||
}
|
||||
if p.Type != "" {
|
||||
data.Set("type", p.Type)
|
||||
}
|
||||
|
||||
resp, err := http.Post(apiURL, "application/x-www-form-urlencoded", strings.NewReader(data.Encode()))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
if resp.StatusCode == 200 && string(body) == "success" {
|
||||
return string(body), nil
|
||||
}
|
||||
|
||||
return string(body), fmt.Errorf("PushMe response error: %s", string(body))
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type PushPlus struct {
|
||||
Token string `json:"token"`
|
||||
Topic string `json:"topic,omitempty"`
|
||||
Template string `json:"template,omitempty"`
|
||||
Channel string `json:"channel,omitempty"`
|
||||
Webhook string `json:"webhook,omitempty"`
|
||||
CallbackUrl string `json:"callbackUrl,omitempty"`
|
||||
To string `json:"to,omitempty"`
|
||||
}
|
||||
|
||||
type pushPlusData struct {
|
||||
PushPlus
|
||||
Title string `json:"title"`
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
type pushPlusResponse struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
Data string `json:"data"`
|
||||
}
|
||||
|
||||
func (p *PushPlus) Request(title, content string) (string, error) {
|
||||
url := "https://www.pushplus.plus/send"
|
||||
|
||||
data := pushPlusData{
|
||||
PushPlus: *p,
|
||||
Title: title,
|
||||
Content: content,
|
||||
}
|
||||
|
||||
body, err := json.Marshal(data)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
resp, err := http.Post(url, "application/json", bytes.NewBuffer(body))
|
||||
if err != nil {
|
||||
// Try old URL if first one fails or as fallback
|
||||
urlOld := "http://pushplus.hxtrip.com/send"
|
||||
resp, err = http.Post(urlOld, "application/json", bytes.NewBuffer(body))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
var res pushPlusResponse
|
||||
if err := json.Unmarshal(respBody, &res); err != nil {
|
||||
return string(respBody), err
|
||||
}
|
||||
|
||||
if res.Code == 200 {
|
||||
return string(respBody), nil
|
||||
}
|
||||
|
||||
return string(respBody), fmt.Errorf("PushPlus error: %s (code: %d)", res.Msg, res.Code)
|
||||
}
|
||||
@@ -0,0 +1,123 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type qywxResponse struct {
|
||||
Code int `json:"errcode"`
|
||||
Msg string `json:"errmsg"`
|
||||
}
|
||||
|
||||
type QyWeiXin struct {
|
||||
AccessToken string
|
||||
}
|
||||
|
||||
func (t *QyWeiXin) Request(msg interface{}) ([]byte, error) {
|
||||
b, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := http.Post(t.getURL(), "application/json", bytes.NewBuffer(b))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
defer func(Body io.ReadCloser) {
|
||||
err := Body.Close()
|
||||
if err != nil {
|
||||
|
||||
}
|
||||
}(resp.Body)
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var r qywxResponse
|
||||
err = json.Unmarshal(body, &r)
|
||||
if err != nil {
|
||||
return body, err
|
||||
}
|
||||
if r.Code != 0 {
|
||||
return body, fmt.Errorf("response error: %s", string(body))
|
||||
}
|
||||
return body, err
|
||||
}
|
||||
|
||||
// SendMessageText Function to send message
|
||||
func (t *QyWeiXin) SendMessageText(text string, at ...string) ([]byte, error) {
|
||||
msg := map[string]interface{}{
|
||||
"msgtype": "text",
|
||||
"text": map[string]interface{}{
|
||||
"content": text,
|
||||
},
|
||||
}
|
||||
|
||||
// 添加@功能
|
||||
// 企业微信支持两种@方式:
|
||||
// 1. mentioned_list: userid列表或"@all"
|
||||
// 2. mentioned_mobile_list: 手机号列表
|
||||
if len(at) > 0 {
|
||||
mentionedList := []string{}
|
||||
mentionedMobileList := []string{}
|
||||
|
||||
for _, item := range at {
|
||||
if item == "@all" || item == "all" {
|
||||
mentionedList = append(mentionedList, "@all")
|
||||
} else if len(item) == 11 && item[0] == '1' {
|
||||
// 判断是否为手机号(简单判断:11位且以1开头)
|
||||
mentionedMobileList = append(mentionedMobileList, item)
|
||||
} else {
|
||||
// 否则当作userid处理
|
||||
mentionedList = append(mentionedList, item)
|
||||
}
|
||||
}
|
||||
|
||||
textContent := msg["text"].(map[string]interface{})
|
||||
if len(mentionedList) > 0 {
|
||||
textContent["mentioned_list"] = mentionedList
|
||||
}
|
||||
if len(mentionedMobileList) > 0 {
|
||||
textContent["mentioned_mobile_list"] = mentionedMobileList
|
||||
}
|
||||
}
|
||||
|
||||
resp, err := t.Request(msg)
|
||||
return resp, err
|
||||
}
|
||||
|
||||
func (t *QyWeiXin) SendMessageMarkdown(title, text string, at ...string) ([]byte, error) {
|
||||
msg := map[string]interface{}{
|
||||
"msgtype": "markdown",
|
||||
"markdown": map[string]interface{}{
|
||||
"content": text,
|
||||
},
|
||||
}
|
||||
|
||||
// 企业微信Markdown消息不支持@功能,但可以在内容中手动添加
|
||||
// 如果需要@功能,建议使用text类型
|
||||
|
||||
resp, err := t.Request(msg)
|
||||
return resp, err
|
||||
}
|
||||
|
||||
func (t *QyWeiXin) getURL() string {
|
||||
url := "https://qyapi.weixin.qq.com/cgi-bin/webhook/send?key=" + t.AccessToken
|
||||
return url
|
||||
}
|
||||
@@ -0,0 +1,193 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/net/proxy"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type telegramResponse struct {
|
||||
Ok bool `json:"ok"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
type Telegram struct {
|
||||
BotToken string
|
||||
ChatID string
|
||||
ApiHost string // 可选的自定义API地址(优先级最高)
|
||||
ProxyURL string // 可选的代理地址,支持 http://、https://、socks5:// 格式
|
||||
}
|
||||
|
||||
func (t *Telegram) Request(params map[string]interface{}) ([]byte, error) {
|
||||
apiURL := t.getAPIURL()
|
||||
|
||||
// 构建请求体
|
||||
data := url.Values{}
|
||||
for key, value := range params {
|
||||
data.Set(key, fmt.Sprintf("%v", value))
|
||||
}
|
||||
|
||||
// 创建 HTTP 客户端
|
||||
client := t.getHTTPClient()
|
||||
|
||||
resp, err := client.Post(apiURL, "application/x-www-form-urlencoded", bytes.NewBufferString(data.Encode()))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
defer func(Body io.ReadCloser) {
|
||||
err := Body.Close()
|
||||
if err != nil {
|
||||
// 忽略关闭错误
|
||||
}
|
||||
}(resp.Body)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var r telegramResponse
|
||||
err = json.Unmarshal(body, &r)
|
||||
if err != nil {
|
||||
return body, err
|
||||
}
|
||||
|
||||
if !r.Ok {
|
||||
return body, fmt.Errorf("telegram api error: %s", r.Description)
|
||||
}
|
||||
|
||||
return body, nil
|
||||
}
|
||||
|
||||
// SendMessageText 发送文本消息
|
||||
func (t *Telegram) SendMessageText(text string) ([]byte, error) {
|
||||
params := map[string]interface{}{
|
||||
"chat_id": t.ChatID,
|
||||
"text": text,
|
||||
"disable_web_page_preview": "true",
|
||||
}
|
||||
|
||||
return t.Request(params)
|
||||
}
|
||||
|
||||
// SendMessageMarkdown 发送Markdown格式消息
|
||||
func (t *Telegram) SendMessageMarkdown(text string) ([]byte, error) {
|
||||
params := map[string]interface{}{
|
||||
"chat_id": t.ChatID,
|
||||
"text": text,
|
||||
"parse_mode": "Markdown",
|
||||
"disable_web_page_preview": "true",
|
||||
}
|
||||
|
||||
return t.Request(params)
|
||||
}
|
||||
|
||||
// SendMessageHTML 发送HTML格式消息
|
||||
func (t *Telegram) SendMessageHTML(text string) ([]byte, error) {
|
||||
params := map[string]interface{}{
|
||||
"chat_id": t.ChatID,
|
||||
"text": text,
|
||||
"parse_mode": "HTML",
|
||||
"disable_web_page_preview": "true",
|
||||
}
|
||||
|
||||
return t.Request(params)
|
||||
}
|
||||
|
||||
func (t *Telegram) getAPIURL() string {
|
||||
// 自定义 API 地址优先级最高
|
||||
if t.ApiHost != "" {
|
||||
return fmt.Sprintf("%s/bot%s/sendMessage", t.ApiHost, t.BotToken)
|
||||
}
|
||||
return fmt.Sprintf("https://api.telegram.org/bot%s/sendMessage", t.BotToken)
|
||||
}
|
||||
|
||||
// getHTTPClient 获取配置了代理的 HTTP 客户端
|
||||
func (t *Telegram) getHTTPClient() *http.Client {
|
||||
client := &http.Client{
|
||||
Timeout: 30 * time.Second,
|
||||
}
|
||||
|
||||
// 如果配置了代理且没有自定义 API 地址,则使用代理
|
||||
// 自定义 API 地址优先级更高,通常用于自建代理服务器
|
||||
if t.ProxyURL != "" && t.ApiHost == "" {
|
||||
proxyURL, err := url.Parse(t.ProxyURL)
|
||||
if err == nil {
|
||||
// 判断是否为 SOCKS5 代理
|
||||
if strings.HasPrefix(strings.ToLower(t.ProxyURL), "socks5://") {
|
||||
// 使用 SOCKS5 代理
|
||||
dialer, err := t.createSOCKS5Dialer(proxyURL)
|
||||
if err == nil {
|
||||
client.Transport = &http.Transport{
|
||||
DialContext: dialer.DialContext,
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// 使用 HTTP/HTTPS 代理
|
||||
client.Transport = &http.Transport{
|
||||
Proxy: http.ProxyURL(proxyURL),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return client
|
||||
}
|
||||
|
||||
// createSOCKS5Dialer 创建 SOCKS5 代理拨号器
|
||||
func (t *Telegram) createSOCKS5Dialer(proxyURL *url.URL) (proxy.ContextDialer, error) {
|
||||
// 解析代理地址
|
||||
host := proxyURL.Host
|
||||
|
||||
// 检查是否有认证信息
|
||||
var auth *proxy.Auth
|
||||
if proxyURL.User != nil {
|
||||
password, _ := proxyURL.User.Password()
|
||||
auth = &proxy.Auth{
|
||||
User: proxyURL.User.Username(),
|
||||
Password: password,
|
||||
}
|
||||
}
|
||||
|
||||
// 创建基础拨号器
|
||||
baseDialer := &net.Dialer{
|
||||
Timeout: 30 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}
|
||||
|
||||
// 创建 SOCKS5 拨号器
|
||||
dialer, err := proxy.SOCKS5("tcp", host, auth, baseDialer)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 转换为 ContextDialer
|
||||
contextDialer, ok := dialer.(proxy.ContextDialer)
|
||||
if !ok {
|
||||
return nil, errors.New("failed to convert to ContextDialer")
|
||||
}
|
||||
|
||||
return contextDialer, nil
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type VoceChat struct {
|
||||
Server string
|
||||
APIKey string
|
||||
TargetType string // "user" or "group"
|
||||
TargetID string
|
||||
}
|
||||
|
||||
func (v *VoceChat) Request(title, content string) ([]byte, error) {
|
||||
if v.Server == "" || v.APIKey == "" || v.TargetID == "" {
|
||||
return nil, fmt.Errorf("vocechat config missing: server, api_key and target_id are required")
|
||||
}
|
||||
|
||||
server := strings.TrimSuffix(v.Server, "/")
|
||||
endpoint := "send_to_user"
|
||||
if v.TargetType == "group" {
|
||||
endpoint = "send_to_group"
|
||||
}
|
||||
|
||||
url := fmt.Sprintf("%s/api/bot/%s/%s", server, endpoint, v.TargetID)
|
||||
|
||||
// Use text/plain for now as requested
|
||||
body := content
|
||||
if title != "" {
|
||||
body = fmt.Sprintf("%s\n\n%s", title, content)
|
||||
}
|
||||
|
||||
req, err := http.NewRequest("POST", url, bytes.NewBuffer([]byte(body)))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "text/plain")
|
||||
req.Header.Set("x-api-key", v.APIKey)
|
||||
|
||||
client := &http.Client{}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return respBody, fmt.Errorf("vocechat response error (status %d): %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
return respBody, nil
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"github.com/silenceper/wechat/v2"
|
||||
"github.com/silenceper/wechat/v2/cache"
|
||||
offConfig "github.com/silenceper/wechat/v2/officialaccount/config"
|
||||
"github.com/silenceper/wechat/v2/officialaccount/message"
|
||||
"github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type WeChatOFAccount struct {
|
||||
AppID string
|
||||
AppSecret string
|
||||
ToUser string
|
||||
TemplateID string
|
||||
URL string
|
||||
}
|
||||
|
||||
// 使用内存缓存进行token的存储
|
||||
var memory = cache.NewMemory()
|
||||
|
||||
func (cw *WeChatOFAccount) Send(title string, content string) (string, error) {
|
||||
wc := wechat.NewWechat()
|
||||
cfg := &offConfig.Config{
|
||||
AppID: cw.AppID,
|
||||
AppSecret: cw.AppSecret,
|
||||
Cache: memory,
|
||||
}
|
||||
officialAccount := wc.GetOfficialAccount(cfg)
|
||||
|
||||
// 获取 Access Token
|
||||
_, err := officialAccount.GetAccessToken()
|
||||
if err != nil {
|
||||
logrus.Errorf("获取access token失败:%s", err)
|
||||
return "", err
|
||||
}
|
||||
|
||||
msgData := make(map[string]*message.TemplateDataItem)
|
||||
msgData["content"] = &message.TemplateDataItem{
|
||||
Value: content,
|
||||
}
|
||||
msgData["title"] = &message.TemplateDataItem{
|
||||
Value: title,
|
||||
//Color: "#173177",
|
||||
}
|
||||
|
||||
// 创建模板消息
|
||||
templateMessage := &message.TemplateMessage{
|
||||
ToUser: cw.ToUser,
|
||||
TemplateID: cw.TemplateID,
|
||||
URL: cw.URL,
|
||||
Data: msgData,
|
||||
}
|
||||
|
||||
// 发送模板消息
|
||||
_, err = officialAccount.GetTemplate().Send(templateMessage)
|
||||
if err != nil {
|
||||
logrus.Errorf("发送模板消息失败: %s", err)
|
||||
return "", err
|
||||
}
|
||||
//logrus.Infof("模板消息发送成功。 消息ID: %d", msgID)
|
||||
return "", nil
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type WxPusher struct {
|
||||
AppToken string `json:"appToken"`
|
||||
Content string `json:"content"`
|
||||
ContentType int `json:"contentType"`
|
||||
TopicIds []int `json:"topicIds,omitempty"`
|
||||
Uids []string `json:"uids,omitempty"`
|
||||
Url string `json:"url,omitempty"`
|
||||
VerifyPayType int `json:"verifyPayType,omitempty"`
|
||||
}
|
||||
|
||||
type wxPusherResponse struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
Data []struct {
|
||||
Uid string `json:"uid"`
|
||||
TopicId int `json:"topicId"`
|
||||
MessageId int `json:"messageId"`
|
||||
Code int `json:"code"`
|
||||
Status string `json:"status"`
|
||||
} `json:"data"`
|
||||
}
|
||||
|
||||
func (w *WxPusher) Send() (string, error) {
|
||||
apiUrl := "https://wxpusher.zjiecode.com/api/send/message"
|
||||
|
||||
body, err := json.Marshal(w)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
resp, err := http.Post(apiUrl, "application/json", bytes.NewBuffer(body))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
var res wxPusherResponse
|
||||
if err := json.Unmarshal(respBody, &res); err != nil {
|
||||
return string(respBody), err
|
||||
}
|
||||
|
||||
if res.Code == 1000 {
|
||||
return string(respBody), nil
|
||||
}
|
||||
|
||||
return string(respBody), fmt.Errorf("WxPusher error: %s (code: %d)", res.Msg, res.Code)
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package channels
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type AliyunSMSChannel struct{ *BaseChannel }
|
||||
|
||||
func NewAliyunSMSChannel() Channel {
|
||||
return &AliyunSMSChannel{NewBaseChannel(ChannelAliyunSMS, []string{FormatTypeText})}
|
||||
}
|
||||
|
||||
func (c *AliyunSMSChannel) Send(config ChannelConfig, msg *Message) (*Result, error) {
|
||||
accessKeyId := config.GetString("access_key_id")
|
||||
accessKeySecret := config.GetString("access_key_secret")
|
||||
signName := config.GetString("sign_name")
|
||||
regionId := config.GetString("region_id")
|
||||
phoneNumber := config.GetString("phone_number")
|
||||
templateCode := config.GetString("template_code")
|
||||
|
||||
if accessKeyId == "" || accessKeySecret == "" || signName == "" {
|
||||
return SendError("aliyun sms config missing: access_key_id, access_key_secret, sign_name are required"), nil
|
||||
}
|
||||
if phoneNumber == "" || templateCode == "" {
|
||||
return SendError("aliyun sms config missing: phone_number, template_code are required"), nil
|
||||
}
|
||||
|
||||
_, formattedContent := c.FormatContent(msg)
|
||||
|
||||
if regionId == "" {
|
||||
regionId = "cn-hangzhou"
|
||||
}
|
||||
|
||||
client, err := createAliyunSMSClient(accessKeyId, accessKeySecret, regionId)
|
||||
if err != nil {
|
||||
return SendError("创建阿里云短信客户端失败: %s", err.Error()), nil
|
||||
}
|
||||
|
||||
result, err := sendAliyunSMS(client, phoneNumber, signName, templateCode, formattedContent, msg.Extra)
|
||||
if err != nil {
|
||||
return ErrorResult("", err), nil
|
||||
}
|
||||
return SuccessResult(result), nil
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package channels
|
||||
|
||||
import "github.com/engigu/taskpool/internal/sdk/message"
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type BarkChannel struct{ *BaseChannel }
|
||||
|
||||
func NewBarkChannel() Channel {
|
||||
return &BarkChannel{NewBaseChannel(ChannelBark, []string{FormatTypeText})}
|
||||
}
|
||||
|
||||
func (c *BarkChannel) Send(config ChannelConfig, msg *Message) (*Result, error) {
|
||||
pushKey := config.GetString("push_key")
|
||||
if pushKey == "" {
|
||||
return SendError("bark config missing: push_key is required"), nil
|
||||
}
|
||||
|
||||
cli := message.Bark{
|
||||
PushKey: pushKey,
|
||||
Archive: config.GetString("archive"),
|
||||
Group: config.GetString("group"),
|
||||
Sound: config.GetString("sound"),
|
||||
Icon: config.GetString("icon"),
|
||||
Level: config.GetString("level"),
|
||||
URL: config.GetString("url"),
|
||||
Key: config.GetString("key"),
|
||||
IV: config.GetString("iv"),
|
||||
Server: config.GetString("server"),
|
||||
Badge: config.GetString("badge"),
|
||||
Copy: config.GetString("copy"),
|
||||
AutoCopy: config.GetString("auto_copy"),
|
||||
ProxyURL: config.GetString("proxy_url"),
|
||||
}
|
||||
|
||||
res, err := cli.Request(msg.Title, msg.Text)
|
||||
if err != nil {
|
||||
return ErrorResult(string(res), err), nil
|
||||
}
|
||||
return SuccessResult(string(res)), nil
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package channels
|
||||
|
||||
import "fmt"
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
// Channel 渠道接口 - SDK 版本,零业务依赖
|
||||
type Channel interface {
|
||||
// GetType 返回渠道类型标识
|
||||
GetType() string
|
||||
// GetSupportedFormats 返回支持的消息格式
|
||||
GetSupportedFormats() []string
|
||||
// Send 发送消息
|
||||
Send(config ChannelConfig, msg *Message) (*Result, error)
|
||||
}
|
||||
|
||||
// BaseChannel 渠道基础实现
|
||||
type BaseChannel struct {
|
||||
channelType string
|
||||
supportedFormats []string
|
||||
}
|
||||
|
||||
func NewBaseChannel(channelType string, supportedFormats []string) *BaseChannel {
|
||||
return &BaseChannel{channelType: channelType, supportedFormats: supportedFormats}
|
||||
}
|
||||
|
||||
func (c *BaseChannel) GetType() string { return c.channelType }
|
||||
func (c *BaseChannel) GetSupportedFormats() []string { return c.supportedFormats }
|
||||
|
||||
// FormatContent 根据渠道支持的格式选择最佳内容
|
||||
func (c *BaseChannel) FormatContent(msg *Message) (formatType string, content string) {
|
||||
for _, ft := range c.supportedFormats {
|
||||
switch ft {
|
||||
case FormatTypeMarkdown:
|
||||
if msg.HasMarkdown() {
|
||||
return FormatTypeMarkdown, msg.Markdown
|
||||
}
|
||||
case FormatTypeHTML:
|
||||
if msg.HasHTML() {
|
||||
return FormatTypeHTML, msg.HTML
|
||||
}
|
||||
case FormatTypeText:
|
||||
if msg.HasText() {
|
||||
return FormatTypeText, msg.Text
|
||||
}
|
||||
}
|
||||
}
|
||||
if msg.HasText() {
|
||||
return FormatTypeText, msg.Text
|
||||
}
|
||||
return FormatTypeText, ""
|
||||
}
|
||||
|
||||
// SuccessResult 创建成功结果
|
||||
func SuccessResult(response string) *Result {
|
||||
return &Result{Success: true, Response: response}
|
||||
}
|
||||
|
||||
// ErrorResult 创建失败结果
|
||||
func ErrorResult(response string, err error) *Result {
|
||||
errMsg := ""
|
||||
if err != nil {
|
||||
errMsg = err.Error()
|
||||
}
|
||||
return &Result{Success: false, Response: response, Error: errMsg}
|
||||
}
|
||||
|
||||
// ErrorResultStr 创建失败结果(字符串错误)
|
||||
func ErrorResultStr(response string, errMsg string) *Result {
|
||||
return &Result{Success: false, Response: response, Error: errMsg}
|
||||
}
|
||||
|
||||
// SendError 发送失败时的格式化错误
|
||||
func SendError(format string, args ...any) *Result {
|
||||
return &Result{Success: false, Error: fmt.Sprintf(format, args...)}
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package channels
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
"github.com/engigu/taskpool/internal/sdk/message"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type CustomChannel struct{ *BaseChannel }
|
||||
|
||||
func NewCustomChannel() Channel {
|
||||
return &CustomChannel{NewBaseChannel(ChannelCustom, []string{FormatTypeText})}
|
||||
}
|
||||
|
||||
func (c *CustomChannel) Send(config ChannelConfig, msg *Message) (*Result, error) {
|
||||
webhook := config.GetString("webhook")
|
||||
body := config.GetString("body")
|
||||
headersStr := config.GetString("headers")
|
||||
|
||||
if webhook == "" {
|
||||
return SendError("custom config missing: webhook is required"), nil
|
||||
}
|
||||
|
||||
var headers map[string]string
|
||||
if headersStr != "" {
|
||||
if err := json.Unmarshal([]byte(headersStr), &headers); err != nil {
|
||||
return SendError("custom config error: headers must be a valid JSON object"), nil
|
||||
}
|
||||
}
|
||||
|
||||
_, formattedContent := c.FormatContent(msg)
|
||||
cli := message.CustomWebhook{}
|
||||
|
||||
// 替换 body 模板中的 TEXT 占位符
|
||||
bodyStr := body
|
||||
if bodyStr != "" {
|
||||
bodyStr = replaceBodyPlaceholder(bodyStr, formattedContent)
|
||||
} else {
|
||||
bodyStr = formattedContent
|
||||
}
|
||||
|
||||
res, err := cli.Request(webhook, bodyStr, headers)
|
||||
if err != nil {
|
||||
return ErrorResult(string(res), err), nil
|
||||
}
|
||||
return SuccessResult(string(res)), nil
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package channels
|
||||
|
||||
import "github.com/engigu/taskpool/internal/sdk/message"
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type DtalkChannel struct{ *BaseChannel }
|
||||
|
||||
func NewDtalkChannel() Channel {
|
||||
return &DtalkChannel{NewBaseChannel(ChannelDtalk, []string{FormatTypeMarkdown, FormatTypeText})}
|
||||
}
|
||||
|
||||
func (c *DtalkChannel) Send(config ChannelConfig, msg *Message) (*Result, error) {
|
||||
accessToken := config.GetString("access_token")
|
||||
secret := config.GetString("secret")
|
||||
|
||||
if accessToken == "" {
|
||||
return SendError("dtalk config missing: access_token is required"), nil
|
||||
}
|
||||
|
||||
contentType, formattedContent := c.FormatContent(msg)
|
||||
atMobiles := msg.GetAtMobiles()
|
||||
if msg.AtAll {
|
||||
atMobiles = append(atMobiles, "all")
|
||||
}
|
||||
|
||||
cli := message.Dtalk{AccessToken: accessToken, Secret: secret}
|
||||
var res []byte
|
||||
var err error
|
||||
|
||||
switch contentType {
|
||||
case FormatTypeText:
|
||||
res, err = cli.SendMessageText(formattedContent, atMobiles...)
|
||||
case FormatTypeMarkdown:
|
||||
res, err = cli.SendMessageMarkdown(msg.Title, formattedContent, atMobiles...)
|
||||
default:
|
||||
return SendError("未知的钉钉发送内容类型:%s", contentType), nil
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return ErrorResult(string(res), err), nil
|
||||
}
|
||||
return SuccessResult(string(res)), nil
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
package channels
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"github.com/engigu/taskpool/internal/sdk/message"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type EmailChannel struct{ *BaseChannel }
|
||||
|
||||
func NewEmailChannel() Channel {
|
||||
return &EmailChannel{NewBaseChannel(ChannelEmail, []string{FormatTypeHTML, FormatTypeText})}
|
||||
}
|
||||
|
||||
func (c *EmailChannel) Send(config ChannelConfig, msg *Message) (*Result, error) {
|
||||
server := config.GetString("server")
|
||||
portStr := config.GetString("port")
|
||||
account := config.GetString("account")
|
||||
passwd := config.GetString("passwd")
|
||||
fromName := config.GetString("from_name")
|
||||
toAccount := config.GetString("to_account")
|
||||
|
||||
if server == "" || account == "" || passwd == "" {
|
||||
return SendError("email config missing: server, account, passwd are required"), nil
|
||||
}
|
||||
if toAccount == "" {
|
||||
return SendError("email config missing: to_account is required"), nil
|
||||
}
|
||||
|
||||
port, _ := strconv.Atoi(portStr)
|
||||
contentType, formattedContent := c.FormatContent(msg)
|
||||
|
||||
var emailer message.EmailMessage
|
||||
emailer.Init(server, port, account, passwd, fromName)
|
||||
|
||||
var errMsg string
|
||||
switch contentType {
|
||||
case FormatTypeText:
|
||||
errMsg = emailer.SendTextMessage(toAccount, msg.Title, formattedContent)
|
||||
case FormatTypeHTML:
|
||||
errMsg = emailer.SendHtmlMessage(toAccount, msg.Title, formattedContent)
|
||||
default:
|
||||
errMsg = fmt.Sprintf("未知的邮件发送内容类型:%s", contentType)
|
||||
}
|
||||
|
||||
if errMsg != "" {
|
||||
return ErrorResultStr("", errMsg), nil
|
||||
}
|
||||
return SuccessResult(""), nil
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
package channels
|
||||
|
||||
import "github.com/engigu/taskpool/internal/sdk/message"
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type FeishuChannel struct{ *BaseChannel }
|
||||
|
||||
func NewFeishuChannel() Channel {
|
||||
return &FeishuChannel{NewBaseChannel(ChannelFeishu, []string{FormatTypeMarkdown, FormatTypeText})}
|
||||
}
|
||||
|
||||
func (c *FeishuChannel) Send(config ChannelConfig, msg *Message) (*Result, error) {
|
||||
accessToken := config.GetString("access_token")
|
||||
secret := config.GetString("secret")
|
||||
|
||||
if accessToken == "" {
|
||||
return SendError("feishu config missing: access_token is required"), nil
|
||||
}
|
||||
|
||||
contentType, formattedContent := c.FormatContent(msg)
|
||||
atMobiles := msg.GetAtMobiles()
|
||||
atUserIds := msg.GetAtUserIds()
|
||||
atList := append(atMobiles, atUserIds...)
|
||||
if msg.AtAll {
|
||||
atList = append(atList, "all")
|
||||
}
|
||||
|
||||
cli := message.Feishu{AccessToken: accessToken, Secret: secret}
|
||||
var res []byte
|
||||
var err error
|
||||
|
||||
switch contentType {
|
||||
case FormatTypeText:
|
||||
res, err = cli.SendMessageText(formattedContent, atList...)
|
||||
case FormatTypeMarkdown:
|
||||
res, err = cli.SendMessageMarkdown(msg.Title, formattedContent, atList...)
|
||||
default:
|
||||
return SendError("未知的飞书发送内容类型:%s", contentType), nil
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return ErrorResult(string(res), err), nil
|
||||
}
|
||||
return SuccessResult(string(res)), nil
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package channels
|
||||
|
||||
import (
|
||||
"github.com/engigu/taskpool/internal/sdk/message"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type GotifyChannel struct{ *BaseChannel }
|
||||
|
||||
func NewGotifyChannel() Channel {
|
||||
return &GotifyChannel{NewBaseChannel(ChannelGotify, []string{FormatTypeText})}
|
||||
}
|
||||
|
||||
func (c *GotifyChannel) Send(config ChannelConfig, msg *Message) (*Result, error) {
|
||||
url := config.GetString("url")
|
||||
token := config.GetString("token")
|
||||
|
||||
if url == "" || token == "" {
|
||||
return SendError("gotify config missing: url and token are required"), nil
|
||||
}
|
||||
|
||||
priority, _ := strconv.Atoi(config.GetString("priority"))
|
||||
cli := message.Gotify{
|
||||
Url: url,
|
||||
Token: token,
|
||||
Priority: priority,
|
||||
}
|
||||
|
||||
res, err := cli.Request(msg.Title, msg.Text)
|
||||
if err != nil {
|
||||
return ErrorResult(string(res), err), nil
|
||||
}
|
||||
return SuccessResult(string(res)), nil
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
package channels
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
openapi "github.com/alibabacloud-go/darabonba-openapi/v2/client"
|
||||
dysmsapi "github.com/alibabacloud-go/dysmsapi-20170525/v4/client"
|
||||
"github.com/alibabacloud-go/tea/tea"
|
||||
)
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
// replaceBodyPlaceholder 替换自定义 webhook body 中的 TEXT 占位符
|
||||
func replaceBodyPlaceholder(body string, content string) string {
|
||||
data, _ := json.Marshal(content)
|
||||
dataStr := strings.Trim(string(data), "\"")
|
||||
return strings.Replace(body, "TEXT", dataStr, -1)
|
||||
}
|
||||
|
||||
// createAliyunSMSClient 创建阿里云短信客户端 (V2.0)
|
||||
func createAliyunSMSClient(accessKeyId, accessKeySecret, regionId string) (*dysmsapi.Client, error) {
|
||||
config := &openapi.Config{
|
||||
AccessKeyId: tea.String(accessKeyId),
|
||||
AccessKeySecret: tea.String(accessKeySecret),
|
||||
RegionId: tea.String(regionId),
|
||||
}
|
||||
// 设置端点,通常为 dysmsapi.aliyuncs.com
|
||||
config.Endpoint = tea.String("dysmsapi.aliyuncs.com")
|
||||
return dysmsapi.NewClient(config)
|
||||
}
|
||||
|
||||
// sendAliyunSMS 发送短信 (V2.0)
|
||||
func sendAliyunSMS(client *dysmsapi.Client, phoneNumber, signName, templateCode, content string, extra map[string]any) (string, error) {
|
||||
templateParam := map[string]interface{}{
|
||||
"content": content,
|
||||
}
|
||||
for k, v := range extra {
|
||||
templateParam[k] = v
|
||||
}
|
||||
templateParamJSON, _ := json.Marshal(templateParam)
|
||||
|
||||
request := &dysmsapi.SendSmsRequest{
|
||||
PhoneNumbers: tea.String(phoneNumber),
|
||||
SignName: tea.String(signName),
|
||||
TemplateCode: tea.String(templateCode),
|
||||
TemplateParam: tea.String(string(templateParamJSON)),
|
||||
}
|
||||
|
||||
response, err := client.SendSms(request)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("发送短信失败: %s", err.Error())
|
||||
}
|
||||
|
||||
if response.Body == nil || tea.StringValue(response.Body.Code) != "OK" {
|
||||
msg := "Unknown Error"
|
||||
code := "Unknown Code"
|
||||
if response.Body != nil {
|
||||
msg = tea.StringValue(response.Body.Message)
|
||||
code = tea.StringValue(response.Body.Code)
|
||||
}
|
||||
return "", fmt.Errorf("发送失败: %s - %s", code, msg)
|
||||
}
|
||||
|
||||
return fmt.Sprintf("RequestId: %s, BizId: %s", tea.StringValue(response.Body.RequestId), tea.StringValue(response.Body.BizId)), nil
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package channels
|
||||
|
||||
import "github.com/engigu/taskpool/internal/sdk/message"
|
||||
|
||||
// Copyright (c) 2026 engigu (TaskPool). All rights reserved.
|
||||
// Use of this source code is governed by the Apache License 2.0.
|
||||
//
|
||||
// 【重要声明 / IMPORTANT NOTICE】
|
||||
// 本代码(包括其架构设计与核心实现)属于任务池(TaskPool)开源项目的一部分。
|
||||
// 任何个人或组织在引用、移植、修改或重新分发此文件中的任何代码时,必须保留本版权声明,
|
||||
// 并在您的衍生作品、文档、软件关于页面或说明文件中显式声明引用自任务池(TaskPool)。
|
||||
//
|
||||
// Anyone referencing, porting, modifying, or redistributing this code must retain this
|
||||
// copyright notice and explicitly state the source: TaskPool (github.com/engigu/taskpool).
|
||||
|
||||
|
||||
type NtfyChannel struct{ *BaseChannel }
|
||||
|
||||
func NewNtfyChannel() Channel {
|
||||
return &NtfyChannel{NewBaseChannel(ChannelNtfy, []string{FormatTypeText})}
|
||||
}
|
||||
|
||||
func (c *NtfyChannel) Send(config ChannelConfig, msg *Message) (*Result, error) {
|
||||
topic := config.GetString("topic")
|
||||
if topic == "" {
|
||||
return SendError("ntfy config missing: topic is required"), nil
|
||||
}
|
||||
|
||||
cli := message.Ntfy{
|
||||
Url: config.GetString("url"),
|
||||
Topic: topic,
|
||||
Priority: config.GetString("priority"),
|
||||
Icon: config.GetString("icon"),
|
||||
Token: config.GetString("token"),
|
||||
Username: config.GetString("username"),
|
||||
Password: config.GetString("password"),
|
||||
Actions: config.GetString("actions"),
|
||||
}
|
||||
|
||||
res, err := cli.Request(msg.Title, msg.Text)
|
||||
if err != nil {
|
||||
return ErrorResult(string(res), err), nil
|
||||
}
|
||||
return SuccessResult(string(res)), nil
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user