Initial commit: TaskPool React panel

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