feat: Initial commit

This commit is contained in:
engigu
2025-12-20 09:30:16 +08:00
commit 362237241e
189 changed files with 12035 additions and 0 deletions
+81
View File
@@ -0,0 +1,81 @@
package controllers
import (
"baihu/internal/middleware"
"baihu/internal/services"
"baihu/internal/utils"
"github.com/gin-gonic/gin"
)
type AuthController struct {
userService *services.UserService
}
func NewAuthController(userService *services.UserService) *AuthController {
return &AuthController{userService: userService}
}
func (ac *AuthController) Login(c *gin.Context) {
var req struct {
Username string `json:"username" binding:"required"`
Password string `json:"password" binding:"required"`
}
if err := c.ShouldBindJSON(&req); err != nil {
utils.BadRequest(c, err.Error())
return
}
user := ac.userService.GetUserByUsername(req.Username)
if user == nil || !ac.userService.ValidatePassword(user, req.Password) {
utils.Unauthorized(c, "用户名或密码错误")
return
}
// 生成 token
token, err := utils.GenerateToken(user.ID, user.Username)
if err != nil {
utils.ServerError(c, "登录失败")
return
}
// 设置 Cookie
middleware.SetAuthCookie(c, token)
utils.Success(c, gin.H{
"user": user.Username,
})
}
func (ac *AuthController) Logout(c *gin.Context) {
middleware.ClearAuthCookie(c)
utils.SuccessMsg(c, "退出成功")
}
func (ac *AuthController) GetCurrentUser(c *gin.Context) {
username, exists := c.Get("username")
if !exists {
utils.Unauthorized(c, "未登录")
return
}
utils.Success(c, gin.H{
"username": username,
})
}
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 := ac.userService.CreateUser(req.Username, req.Email, req.Password, "user")
utils.Success(c, user)
}
@@ -0,0 +1,51 @@
package controllers
import (
"baihu/internal/database"
"baihu/internal/models"
"baihu/internal/services"
"baihu/internal/utils"
"github.com/gin-gonic/gin"
)
type DashboardController struct {
cronService *services.CronService
executorService *services.ExecutorService
}
func NewDashboardController(cronService *services.CronService, executorService *services.ExecutorService) *DashboardController {
return &DashboardController{
cronService: cronService,
executorService: executorService,
}
}
type StatsResponse struct {
Tasks int64 `json:"tasks"`
Scripts int64 `json:"scripts"`
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, scriptCount, envCount, logCount int64
database.DB.Model(&models.Task{}).Count(&taskCount)
database.DB.Model(&models.Script{}).Count(&scriptCount)
database.DB.Model(&models.EnvironmentVariable{}).Count(&envCount)
database.DB.Model(&models.TaskLog{}).Count(&logCount)
stats := StatsResponse{
Tasks: taskCount,
Scripts: scriptCount,
Envs: envCount,
Logs: logCount,
Scheduled: dc.cronService.GetScheduledCount(),
Running: dc.executorService.GetRunningCount(),
}
utils.Success(c, stats)
}
+103
View File
@@ -0,0 +1,103 @@
package controllers
import (
"strconv"
"baihu/internal/services"
"baihu/internal/utils"
"github.com/gin-gonic/gin"
)
type EnvController struct {
envService *services.EnvService
}
func NewEnvController(envService *services.EnvService) *EnvController {
return &EnvController{envService: envService}
}
func (ec *EnvController) CreateEnvVar(c *gin.Context) {
userID := 1
var req struct {
Name string `json:"name" binding:"required"`
Value string `json:"value" binding:"required"`
Remark string `json:"remark"`
}
if err := c.ShouldBindJSON(&req); err != nil {
utils.BadRequest(c, err.Error())
return
}
envVar := ec.envService.CreateEnvVar(req.Name, req.Value, req.Remark, userID)
utils.Success(c, envVar)
}
func (ec *EnvController) GetEnvVars(c *gin.Context) {
userID := 1
p := utils.ParsePagination(c)
name := c.DefaultQuery("name", "")
envVars, total := ec.envService.GetEnvVarsWithPagination(userID, name, p.Page, p.PageSize)
utils.PaginatedResponse(c, envVars, total, p)
}
func (ec *EnvController) GetEnvVar(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的环境变量ID")
return
}
envVar := ec.envService.GetEnvVarByID(id)
if envVar == nil {
utils.NotFound(c, "环境变量不存在")
return
}
utils.Success(c, envVar)
}
func (ec *EnvController) UpdateEnvVar(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的环境变量ID")
return
}
var req struct {
Name string `json:"name"`
Value string `json:"value"`
Remark string `json:"remark"`
}
if err := c.ShouldBindJSON(&req); err != nil {
utils.BadRequest(c, err.Error())
return
}
envVar := ec.envService.UpdateEnvVar(id, req.Name, req.Value, req.Remark)
if envVar == nil {
utils.NotFound(c, "环境变量不存在")
return
}
utils.Success(c, envVar)
}
func (ec *EnvController) DeleteEnvVar(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的环境变量ID")
return
}
success := ec.envService.DeleteEnvVar(id)
if !success {
utils.NotFound(c, "环境变量不存在")
return
}
utils.SuccessMsg(c, "删除成功")
}
@@ -0,0 +1,55 @@
package controllers
import (
"strconv"
"baihu/internal/services"
"baihu/internal/utils"
"github.com/gin-gonic/gin"
)
type ExecutorController struct {
executorService *services.ExecutorService
}
func NewExecutorController(executorService *services.ExecutorService) *ExecutorController {
return &ExecutorController{executorService: executorService}
}
func (ec *ExecutorController) ExecuteTask(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的任务ID")
return
}
result := ec.executorService.ExecuteTask(id)
utils.Success(c, result)
}
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, result)
}
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, results)
}
+344
View File
@@ -0,0 +1,344 @@
package controllers
import (
"io/fs"
"os"
"path/filepath"
"strings"
"baihu/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"`
Children []*FileNode `json:"children,omitempty"`
}
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
}
relPath, _ := filepath.Rel(fc.workDir, path)
parts := strings.Split(relPath, string(filepath.Separator))
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,
}
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 := filepath.Join(fc.workDir, filepath.Clean(filePath))
if !strings.HasPrefix(fullPath, fc.workDir) {
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 := filepath.Join(fc.workDir, filepath.Clean(req.Path))
if !strings.HasPrefix(fullPath, fc.workDir) {
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 := filepath.Join(fc.workDir, filepath.Clean(req.Path))
if !strings.HasPrefix(fullPath, fc.workDir) {
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 := filepath.Join(fc.workDir, filepath.Clean(req.Path))
if !strings.HasPrefix(fullPath, fc.workDir) {
utils.Forbidden(c, "访问被拒绝")
return
}
if err := os.RemoveAll(fullPath); err != nil {
utils.ServerError(c, err.Error())
return
}
utils.SuccessMsg(c, "删除成功")
}
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
}
oldFull := filepath.Join(fc.workDir, filepath.Clean(req.OldPath))
newFull := filepath.Join(fc.workDir, filepath.Clean(req.NewPath))
if !strings.HasPrefix(oldFull, fc.workDir) || !strings.HasPrefix(newFull, fc.workDir) {
utils.Forbidden(c, "访问被拒绝")
return
}
// 确保目标目录存在
os.MkdirAll(filepath.Dir(newFull), 0755)
if err := os.Rename(oldFull, newFull); err != nil {
utils.ServerError(c, err.Error())
return
}
utils.SuccessMsg(c, "移动成功")
}
// UploadArchive handles archive file upload and extraction
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 := fc.workDir
if targetDir != "" {
extractDir = filepath.Join(fc.workDir, filepath.Clean(targetDir))
if !strings.HasPrefix(extractDir, fc.workDir) {
utils.Forbidden(c, "访问被拒绝")
return
}
}
os.MkdirAll(extractDir, 0755)
// 保存临时文件
tempFile := filepath.Join(os.TempDir(), 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 handles multiple file uploads
func (fc *FileController) UploadFiles(c *gin.Context) {
targetDir := c.PostForm("path")
// 确定目标目录
destDir := fc.workDir
if targetDir != "" {
destDir = filepath.Join(fc.workDir, filepath.Clean(targetDir))
if !strings.HasPrefix(destDir, fc.workDir) {
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 := file.Filename
if i < len(paths) && paths[i] != "" {
relPath = paths[i]
}
// 构建完整路径
fullPath := filepath.Join(destDir, filepath.Clean(relPath))
// 安全检查
if !strings.HasPrefix(fullPath, fc.workDir) {
continue
}
// 确保父目录存在
os.MkdirAll(filepath.Dir(fullPath), 0755)
// 保存文件
if err := c.SaveUploadedFile(file, fullPath); err != nil {
utils.ServerError(c, "保存文件失败: "+err.Error())
return
}
}
utils.SuccessMsg(c, "上传成功")
}
+107
View File
@@ -0,0 +1,107 @@
package controllers
import (
"strconv"
"baihu/internal/database"
"baihu/internal/models"
"baihu/internal/utils"
"github.com/gin-gonic/gin"
)
type LogController struct{}
func NewLogController() *LogController {
return &LogController{}
}
type TaskLogResponse struct {
ID uint `json:"id"`
TaskID uint `json:"task_id"`
TaskName string `json:"task_name"`
Command string `json:"command"`
Status string `json:"status"`
Duration int64 `json:"duration"`
CreatedAt models.LocalTime `json:"created_at"`
}
func (lc *LogController) GetLogs(c *gin.Context) {
p := utils.ParsePagination(c)
taskID, _ := strconv.Atoi(c.DefaultQuery("task_id", "0"))
taskName := c.DefaultQuery("task_name", "")
var logs []models.TaskLog
var total int64
query := database.DB.Model(&models.TaskLog{})
if taskID > 0 {
query = query.Where("task_id = ?", taskID)
}
// 按任务名称过滤
if taskName != "" {
var taskIDs []uint
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, []TaskLogResponse{}, 0, p)
return
}
}
query.Count(&total)
query.Order("id DESC").Offset(p.Offset()).Limit(p.PageSize).Find(&logs)
taskIDList := make([]uint, 0)
for _, log := range logs {
taskIDList = append(taskIDList, log.TaskID)
}
var tasks []models.Task
database.DB.Where("id IN ?", taskIDList).Find(&tasks)
taskMap := make(map[uint]string)
for _, t := range tasks {
taskMap[t.ID] = t.Name
}
result := make([]TaskLogResponse, len(logs))
for i, log := range logs {
result[i] = TaskLogResponse{
ID: log.ID,
TaskID: log.TaskID,
TaskName: taskMap[log.TaskID],
Command: log.Command,
Status: log.Status,
Duration: log.Duration,
CreatedAt: log.CreatedAt,
}
}
utils.PaginatedResponse(c, result, total, p)
}
func (lc *LogController) GetLogDetail(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的日志ID")
return
}
var log models.TaskLog
if err := database.DB.First(&log, id).Error; err != nil {
utils.NotFound(c, "日志不存在")
return
}
utils.Success(c, gin.H{
"id": log.ID,
"task_id": log.TaskID,
"command": log.Command,
"output": log.Output,
"status": log.Status,
"duration": log.Duration,
"created_at": log.CreatedAt,
})
}
+99
View File
@@ -0,0 +1,99 @@
package controllers
import (
"strconv"
"baihu/internal/services"
"baihu/internal/utils"
"github.com/gin-gonic/gin"
)
type ScriptController struct {
scriptService *services.ScriptService
}
func NewScriptController(scriptService *services.ScriptService) *ScriptController {
return &ScriptController{scriptService: scriptService}
}
func (sc *ScriptController) CreateScript(c *gin.Context) {
userID := 1
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, script)
}
func (sc *ScriptController) GetScripts(c *gin.Context) {
userID := 1
scripts := sc.scriptService.GetScriptsByUserID(userID)
utils.Success(c, scripts)
}
func (sc *ScriptController) GetScript(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的脚本ID")
return
}
script := sc.scriptService.GetScriptByID(id)
if script == nil {
utils.NotFound(c, "脚本不存在")
return
}
utils.Success(c, script)
}
func (sc *ScriptController) UpdateScript(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
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, script)
}
func (sc *ScriptController) DeleteScript(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的脚本ID")
return
}
success := sc.scriptService.DeleteScript(id)
if !success {
utils.NotFound(c, "脚本不存在")
return
}
utils.SuccessMsg(c, "删除成功")
}
+149
View File
@@ -0,0 +1,149 @@
package controllers
import (
"baihu/internal/constant"
"baihu/internal/database"
"baihu/internal/models"
"baihu/internal/services"
"baihu/internal/utils"
"fmt"
"runtime"
"time"
"github.com/gin-gonic/gin"
)
type SettingsController struct {
userService *services.UserService
}
func NewSettingsController(userService *services.UserService) *SettingsController {
return &SettingsController{userService: userService}
}
// ChangePassword 修改密码
func (sc *SettingsController) ChangePassword(c *gin.Context) {
var req struct {
OldPassword string `json:"old_password" binding:"required"`
NewPassword string `json:"new_password" binding:"required,min=6"`
}
if err := c.ShouldBindJSON(&req); err != nil {
utils.BadRequest(c, "参数错误")
return
}
// 暂时使用固定用户名 admin
user := sc.userService.GetUserByUsername("admin")
if user == nil {
utils.NotFound(c, "用户不存在")
return
}
if !sc.userService.ValidatePassword(user, req.OldPassword) {
utils.BadRequest(c, "原密码错误")
return
}
if err := sc.userService.UpdatePassword(user.ID, req.NewPassword); err != nil {
utils.ServerError(c, "修改密码失败")
return
}
utils.SuccessMsg(c, "密码修改成功")
}
// CleanLogs 清理日志
func (sc *SettingsController) CleanLogs(c *gin.Context) {
var req struct {
Days int `json:"days" binding:"required,min=1"`
}
if err := c.ShouldBindJSON(&req); err != nil {
utils.BadRequest(c, "参数错误")
return
}
cutoff := time.Now().AddDate(0, 0, -req.Days)
result := database.DB.Where("created_at < ?", cutoff).Delete(&models.TaskLog{})
utils.Success(c, gin.H{
"deleted": result.RowsAffected,
})
}
// GetSiteSettings 获取站点设置
func (sc *SettingsController) GetSiteSettings(c *gin.Context) {
config := services.Config
if config == nil {
utils.Success(c, gin.H{
"site_name": "白虎面板",
"port": 8080,
})
return
}
utils.Success(c, gin.H{
"site_name": config.Server.SiteName,
"port": config.Server.Port,
})
}
// 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)
// 内存使用
var m runtime.MemStats
runtime.ReadMemStats(&m)
memUsage := formatBytes(m.Alloc)
// 运行时间
uptime := formatDuration(time.Since(constant.StartTime))
utils.Success(c, gin.H{
"version": constant.Version,
"build_time": constant.BuildTime,
"mem_usage": memUsage,
"uptime": uptime,
"task_count": taskCount,
"log_count": logCount,
"env_count": envCount,
})
}
// 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)
}
+128
View File
@@ -0,0 +1,128 @@
package controllers
import (
"strconv"
"baihu/internal/services"
"baihu/internal/utils"
"github.com/gin-gonic/gin"
)
type TaskController struct {
taskService *services.TaskService
cronService *services.CronService
}
func NewTaskController(taskService *services.TaskService, cronService *services.CronService) *TaskController {
return &TaskController{
taskService: taskService,
cronService: cronService,
}
}
func (tc *TaskController) CreateTask(c *gin.Context) {
var req struct {
Name string `json:"name" binding:"required"`
Command string `json:"command" binding:"required"`
Schedule string `json:"schedule" binding:"required"`
}
if err := c.ShouldBindJSON(&req); err != nil {
utils.BadRequest(c, err.Error())
return
}
if err := tc.cronService.ValidateCron(req.Schedule); err != nil {
utils.BadRequest(c, "无效的cron表达式: "+err.Error())
return
}
task := tc.taskService.CreateTask(req.Name, req.Command, req.Schedule)
tc.cronService.AddTask(task)
utils.Success(c, task)
}
func (tc *TaskController) GetTasks(c *gin.Context) {
p := utils.ParsePagination(c)
name := c.DefaultQuery("name", "")
tasks, total := tc.taskService.GetTasksWithPagination(p.Page, p.PageSize, name)
utils.PaginatedResponse(c, tasks, total, p)
}
func (tc *TaskController) GetTask(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的任务ID")
return
}
task := tc.taskService.GetTaskByID(id)
if task == nil {
utils.NotFound(c, "任务不存在")
return
}
utils.Success(c, task)
}
func (tc *TaskController) UpdateTask(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的任务ID")
return
}
var req struct {
Name string `json:"name"`
Command string `json:"command"`
Schedule string `json:"schedule"`
Enabled bool `json:"enabled"`
}
if err := c.ShouldBindJSON(&req); err != nil {
utils.BadRequest(c, err.Error())
return
}
if req.Schedule != "" {
if err := tc.cronService.ValidateCron(req.Schedule); err != nil {
utils.BadRequest(c, "无效的cron表达式: "+err.Error())
return
}
}
task := tc.taskService.UpdateTask(id, req.Name, req.Command, req.Schedule, req.Enabled)
if task == nil {
utils.NotFound(c, "任务不存在")
return
}
if task.Enabled {
tc.cronService.AddTask(task)
} else {
tc.cronService.RemoveTask(task.ID)
}
utils.Success(c, task)
}
func (tc *TaskController) DeleteTask(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
utils.BadRequest(c, "无效的任务ID")
return
}
tc.cronService.RemoveTask(uint(id))
success := tc.taskService.DeleteTask(id)
if !success {
utils.NotFound(c, "任务不存在")
return
}
utils.SuccessMsg(c, "删除成功")
}
+249
View File
@@ -0,0 +1,249 @@
package controllers
import (
"bufio"
"io"
"net/http"
"os"
"path/filepath"
"runtime"
"sync"
"unicode/utf8"
"baihu/internal/constant"
"baihu/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{}
func NewTerminalController() *TerminalController {
return &TerminalController{}
}
var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
return true
},
}
// 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()
// Windows 使用 pipe 模式,Unix 使用 PTY 模式
if runtime.GOOS == "windows" {
tc.handlePipeMode(conn)
} else {
tc.handlePtyMode(conn)
}
}
// handlePtyMode 使用 PTY 处理终端(Unix/macOS
func (tc *TerminalController) handlePtyMode(conn *websocket.Conn) {
// 发送 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 = append(os.Environ(), "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.WriteMessage(websocket.TextMessage, data)
}
wg.Add(1)
go func() {
defer wg.Done()
buf := make([]byte, 4096)
for {
n, err := ptmx.Read(buf)
if err != nil {
return
}
if n > 0 {
text := toUTF8(buf[:n])
writeMessage([]byte(text))
}
}
}()
for {
_, message, err := conn.ReadMessage()
if err != nil {
break
}
if _, err := ptmx.Write(message); err != nil {
break
}
}
cmd.Process.Kill()
cmd.Wait()
wg.Wait()
}
// handlePipeMode 使用 pipe 处理终端(Windows
func (tc *TerminalController) handlePipeMode(conn *websocket.Conn) {
// 发送 pipe 模式标识
conn.WriteMessage(websocket.TextMessage, []byte("__PIPE_MODE__"))
cmd := utils.NewShellCmd()
if absDir, err := filepath.Abs(constant.ScriptsWorkDir); err == nil {
cmd.Dir = absDir
}
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.WriteMessage(websocket.TextMessage, data)
}
readOutput := func(reader io.Reader) {
defer wg.Done()
defer func() { recover() }()
buf := make([]byte, 4096)
for {
n, err := reader.Read(buf)
if err != nil {
return
}
if n > 0 {
text := toUTF8(buf[:n])
writeMessage([]byte(text))
}
}
}
wg.Add(2)
go readOutput(stdout)
go readOutput(stderr)
for {
_, message, err := conn.ReadMessage()
if err != nil {
break
}
if _, err := stdin.Write(message); err != nil {
break
}
}
stdin.Close()
cmd.Process.Kill()
cmd.Wait()
wg.Wait()
}
// ExecuteShellCommand 执行单个命令并返回结果
func (tc *TerminalController) ExecuteShellCommand(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
}
cmd := utils.NewShellCommandCmd(req.Command)
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),
})
}