feat: Initial commit
This commit is contained in:
@@ -0,0 +1,76 @@
|
||||
package bootstrap
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"baihu/internal/constant"
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/logger"
|
||||
"baihu/internal/router"
|
||||
"baihu/internal/services"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type App struct {
|
||||
Config *services.AppConfig
|
||||
Router *gin.Engine
|
||||
}
|
||||
|
||||
func New() *App {
|
||||
app := &App{}
|
||||
app.initConfig()
|
||||
app.initDatabase()
|
||||
app.initRouter()
|
||||
return app
|
||||
}
|
||||
|
||||
func (a *App) initConfig() {
|
||||
cfg, err := services.LoadConfig(constant.ConfigPath)
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
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,
|
||||
}
|
||||
|
||||
if err := database.Init(dbCfg); err != nil {
|
||||
logger.Fatalf("Failed to init database: %v", err)
|
||||
}
|
||||
|
||||
if err := database.Migrate(); err != nil {
|
||||
logger.Fatalf("Failed to migrate database: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) initRouter() {
|
||||
ctrls := router.RegisterControllers()
|
||||
a.Router = router.Setup(ctrls)
|
||||
}
|
||||
|
||||
func (a *App) Run() {
|
||||
addr := fmt.Sprintf("%s:%d", a.Config.Server.Host, a.Config.Server.Port)
|
||||
logger.Infof("Starting server on %s", addr)
|
||||
a.Router.Run(addr)
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package constant
|
||||
|
||||
const (
|
||||
|
||||
// ConfigPath 配置文件路径
|
||||
ConfigPath = "configs/config.json"
|
||||
|
||||
// DataDir 数据目录
|
||||
DataDir = "./data"
|
||||
|
||||
// WebDistDir 前端构建目录
|
||||
WebDistDir = "./web/dist"
|
||||
|
||||
// DefaultRole 默认用户角色
|
||||
DefaultRole = "user"
|
||||
|
||||
// AdminRole 管理员角色
|
||||
AdminRole = "admin"
|
||||
|
||||
// DefaultTablePrefix 默认表前缀
|
||||
DefaultTablePrefix = "baihu_"
|
||||
|
||||
// ScriptsWorkDir 脚本工作目录
|
||||
ScriptsWorkDir = "./data/scripts"
|
||||
|
||||
// DefaultPageSize 默认分页大小
|
||||
DefaultPageSize = 10
|
||||
|
||||
// CookieName Cookie 名称
|
||||
CookieName = "BHToken"
|
||||
|
||||
// TokenExpireDays Token 过期天数
|
||||
TokenExpireDays = 7
|
||||
// CookieMaxAge Cookie 有效期(秒)7天
|
||||
CookieMaxAge = 86400 * TokenExpireDays
|
||||
|
||||
// DefaultJWTSecret 默认 JWT 密钥
|
||||
DefaultJWTSecret = "baihu-default-secret-key"
|
||||
|
||||
// DefaultTaskTimeout 默认任务超时时间(分钟)
|
||||
DefaultTaskTimeout = 30
|
||||
)
|
||||
|
||||
// TablePrefix 表前缀,可在运行时设置
|
||||
var TablePrefix = DefaultTablePrefix
|
||||
|
||||
// JWTSecret JWT 密钥,可通过配置文件设置
|
||||
var JWTSecret = DefaultJWTSecret
|
||||
@@ -0,0 +1,12 @@
|
||||
package constant
|
||||
|
||||
import "time"
|
||||
|
||||
// 构建时注入的变量
|
||||
var (
|
||||
Version = "dev"
|
||||
BuildTime = "unknown"
|
||||
)
|
||||
|
||||
// 程序启动时间
|
||||
var StartTime = time.Now()
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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, "上传成功")
|
||||
}
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
@@ -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, "删除成功")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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, "删除成功")
|
||||
}
|
||||
@@ -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),
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"baihu/internal/logger"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/driver/mysql"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
gormlogger "gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
var DB *gorm.DB
|
||||
|
||||
type Config struct {
|
||||
Type string // sqlite, mysql, postgres
|
||||
Host string
|
||||
Port int
|
||||
User string
|
||||
Password string
|
||||
DBName string
|
||||
Path string // for sqlite
|
||||
}
|
||||
|
||||
func Init(cfg *Config) error {
|
||||
// 设置东八区时区
|
||||
loc, err := time.LoadLocation("Asia/Shanghai")
|
||||
if err != nil {
|
||||
logger.Warnf("Failed to load timezone, using UTC: %v", err)
|
||||
loc = time.UTC
|
||||
}
|
||||
time.Local = loc
|
||||
|
||||
var dialector gorm.Dialector
|
||||
|
||||
switch cfg.Type {
|
||||
case "sqlite":
|
||||
dialector = sqlite.Open(cfg.Path)
|
||||
case "mysql":
|
||||
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)
|
||||
dialector = mysql.Open(dsn)
|
||||
case "postgres":
|
||||
dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s sslmode=disable TimeZone=Asia/Shanghai",
|
||||
cfg.Host, cfg.Port, cfg.User, cfg.Password, cfg.DBName)
|
||||
dialector = postgres.Open(dsn)
|
||||
default:
|
||||
return fmt.Errorf("unsupported database type: %s", cfg.Type)
|
||||
}
|
||||
|
||||
DB, err = gorm.Open(dialector, &gorm.Config{
|
||||
Logger: gormlogger.Default.LogMode(gormlogger.Warn),
|
||||
NowFunc: func() time.Time {
|
||||
return time.Now().In(loc)
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to connect database: %w", err)
|
||||
}
|
||||
|
||||
logger.Infof("Connected to %s database with Asia/Shanghai timezone", cfg.Type)
|
||||
return nil
|
||||
}
|
||||
|
||||
func AutoMigrate(models ...interface{}) error {
|
||||
return DB.AutoMigrate(models...)
|
||||
}
|
||||
|
||||
func GetDB() *gorm.DB {
|
||||
return DB
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"baihu/internal/models"
|
||||
)
|
||||
|
||||
func Migrate() error {
|
||||
return AutoMigrate(
|
||||
&models.User{},
|
||||
&models.Task{},
|
||||
&models.TaskLog{},
|
||||
&models.Script{},
|
||||
&models.EnvironmentVariable{},
|
||||
&models.Setting{},
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package logger
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
var Log *logrus.Logger
|
||||
|
||||
func init() {
|
||||
Log = logrus.New()
|
||||
|
||||
// 设置日志格式
|
||||
Log.SetFormatter(&logrus.TextFormatter{
|
||||
FullTimestamp: true,
|
||||
TimestampFormat: "2006-01-02 15:04:05",
|
||||
ForceColors: true,
|
||||
})
|
||||
|
||||
// 设置日志级别
|
||||
Log.SetLevel(logrus.InfoLevel)
|
||||
|
||||
// 输出到标准输出
|
||||
Log.SetOutput(os.Stdout)
|
||||
}
|
||||
|
||||
// SetupFileOutput 设置文件输出
|
||||
func SetupFileOutput(logDir string) error {
|
||||
if err := os.MkdirAll(logDir, 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
logFile := filepath.Join(logDir, time.Now().Format("2006-01-02")+".log")
|
||||
file, err := os.OpenFile(logFile, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0666)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
Log.SetOutput(file)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetLevel 设置日志级别
|
||||
func SetLevel(level string) {
|
||||
switch level {
|
||||
case "debug":
|
||||
Log.SetLevel(logrus.DebugLevel)
|
||||
case "info":
|
||||
Log.SetLevel(logrus.InfoLevel)
|
||||
case "warn":
|
||||
Log.SetLevel(logrus.WarnLevel)
|
||||
case "error":
|
||||
Log.SetLevel(logrus.ErrorLevel)
|
||||
default:
|
||||
Log.SetLevel(logrus.InfoLevel)
|
||||
}
|
||||
}
|
||||
|
||||
// 便捷方法
|
||||
func Debug(args ...interface{}) { Log.Debug(args...) }
|
||||
func Info(args ...interface{}) { Log.Info(args...) }
|
||||
func Warn(args ...interface{}) { Log.Warn(args...) }
|
||||
func Error(args ...interface{}) { Log.Error(args...) }
|
||||
func Fatal(args ...interface{}) { Log.Fatal(args...) }
|
||||
|
||||
func Debugf(format string, args ...interface{}) { Log.Debugf(format, args...) }
|
||||
func Infof(format string, args ...interface{}) { Log.Infof(format, args...) }
|
||||
func Warnf(format string, args ...interface{}) { Log.Warnf(format, args...) }
|
||||
func Errorf(format string, args ...interface{}) { Log.Errorf(format, args...) }
|
||||
func Fatalf(format string, args ...interface{}) { Log.Fatalf(format, args...) }
|
||||
|
||||
// WithField 带字段的日志
|
||||
func WithField(key string, value interface{}) *logrus.Entry {
|
||||
return Log.WithField(key, value)
|
||||
}
|
||||
|
||||
// WithFields 带多个字段的日志
|
||||
func WithFields(fields logrus.Fields) *logrus.Entry {
|
||||
return Log.WithFields(fields)
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"baihu/internal/constant"
|
||||
"baihu/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// AuthRequired 认证中间件
|
||||
func AuthRequired() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token, err := c.Cookie(constant.CookieName)
|
||||
if err != nil || token == "" {
|
||||
utils.Unauthorized(c, "请先登录")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 验证 token
|
||||
userID, username, err := utils.ParseToken(token)
|
||||
if err != nil {
|
||||
utils.Unauthorized(c, "登录已过期,请重新登录")
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
// 将用户信息存入上下文
|
||||
c.Set("userID", userID)
|
||||
c.Set("username", username)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// SetAuthCookie 设置认证 Cookie
|
||||
func SetAuthCookie(c *gin.Context, token string) {
|
||||
c.SetCookie(constant.CookieName, token, constant.CookieMaxAge, "/", "", false, true)
|
||||
}
|
||||
|
||||
// ClearAuthCookie 清除认证 Cookie
|
||||
func ClearAuthCookie(c *gin.Context) {
|
||||
c.SetCookie(constant.CookieName, "", -1, "/", "", false, true)
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"baihu/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
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("%3d | %13v | %15s | %-7s %s",
|
||||
status, latency, clientIP, method, path)
|
||||
|
||||
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("Panic recovered: %v | path: %s", err, c.Request.URL.Path)
|
||||
c.AbortWithStatus(500)
|
||||
}
|
||||
}()
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"baihu/internal/constant"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// EnvironmentVariable represents an environment variable
|
||||
type EnvironmentVariable struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Name string `json:"name" gorm:"size:255;not null"`
|
||||
Value string `json:"value" gorm:"type:text"`
|
||||
Remark string `json:"remark" gorm:"size:500"`
|
||||
UserID uint `json:"user_id" gorm:"index"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
DeletedAt gorm.DeletedAt `json:"-" gorm:"index"`
|
||||
}
|
||||
|
||||
func (EnvironmentVariable) TableName() string {
|
||||
return constant.TablePrefix + "envs"
|
||||
}
|
||||
|
||||
// Script represents a script file
|
||||
type Script struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Name string `json:"name" gorm:"size:255;not null"`
|
||||
Content string `json:"content" gorm:"type:text"`
|
||||
UserID uint `json:"user_id" gorm:"index"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
DeletedAt gorm.DeletedAt `json:"-" gorm:"index"`
|
||||
}
|
||||
|
||||
func (Script) TableName() string {
|
||||
return constant.TablePrefix + "scripts"
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"baihu/internal/constant"
|
||||
)
|
||||
|
||||
// Setting 系统设置
|
||||
type Setting struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
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 string `json:"value" gorm:"type:text"`
|
||||
}
|
||||
|
||||
func (Setting) TableName() string {
|
||||
return constant.TablePrefix + "settings"
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"baihu/internal/constant"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Task represents a scheduled task
|
||||
type Task struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Name string `json:"name" gorm:"size:255;not null"`
|
||||
Command string `json:"command" gorm:"type:text;not null"`
|
||||
Schedule string `json:"schedule" gorm:"size:100"` // cron expression
|
||||
Timeout int `json:"timeout" gorm:"default:30"` // 超时时间(分钟),默认30分钟
|
||||
Enabled bool `json:"enabled" gorm:"default:true"`
|
||||
LastRun *LocalTime `json:"last_run"`
|
||||
NextRun *LocalTime `json:"next_run"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
DeletedAt gorm.DeletedAt `json:"-" gorm:"index"`
|
||||
}
|
||||
|
||||
func (Task) TableName() string {
|
||||
return constant.TablePrefix + "tasks"
|
||||
}
|
||||
|
||||
// TaskLog represents a log entry for task execution
|
||||
type TaskLog struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
TaskID uint `json:"task_id" gorm:"index"`
|
||||
Command string `json:"command" gorm:"type:text"`
|
||||
Output string `json:"-" gorm:"type:longtext"` // gzip+base64 compressed
|
||||
Status string `json:"status" gorm:"size:20"` // success, failed
|
||||
Duration int64 `json:"duration"` // milliseconds
|
||||
ExitCode int `json:"exit_code"`
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
}
|
||||
|
||||
func (TaskLog) TableName() string {
|
||||
return constant.TablePrefix + "task_logs"
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
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
|
||||
}
|
||||
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(time.Now())
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"baihu/internal/constant"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// User represents a system user
|
||||
type User struct {
|
||||
ID uint `json:"id" gorm:"primaryKey"`
|
||||
Username string `json:"username" gorm:"size:100;uniqueIndex;not null"`
|
||||
Password string `json:"-" gorm:"size:255;not null"`
|
||||
Email string `json:"email" gorm:"size:255"`
|
||||
Role string `json:"role" gorm:"size:20;default:user"` // admin, user
|
||||
CreatedAt LocalTime `json:"created_at"`
|
||||
UpdatedAt LocalTime `json:"updated_at"`
|
||||
DeletedAt gorm.DeletedAt `json:"-" gorm:"index"`
|
||||
}
|
||||
|
||||
func (User) TableName() string {
|
||||
return constant.TablePrefix + "users"
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"baihu/internal/constant"
|
||||
"baihu/internal/controllers"
|
||||
"baihu/internal/services"
|
||||
)
|
||||
|
||||
var cronService *services.CronService
|
||||
|
||||
func RegisterControllers() *Controllers {
|
||||
// Initialize services
|
||||
taskService := services.NewTaskService()
|
||||
userService := services.NewUserService()
|
||||
envService := services.NewEnvService()
|
||||
scriptService := services.NewScriptService()
|
||||
executorService := services.NewExecutorService(taskService)
|
||||
settingsService := services.NewSettingsService()
|
||||
|
||||
// 执行系统初始化
|
||||
initService := services.NewInitService(settingsService, userService)
|
||||
initService.Initialize()
|
||||
|
||||
// Initialize cron service
|
||||
cronService = services.NewCronService(taskService, executorService)
|
||||
cronService.Start()
|
||||
|
||||
// Initialize and return controllers
|
||||
return &Controllers{
|
||||
Task: controllers.NewTaskController(taskService, cronService),
|
||||
Auth: controllers.NewAuthController(userService),
|
||||
Env: controllers.NewEnvController(envService),
|
||||
Script: controllers.NewScriptController(scriptService),
|
||||
Executor: controllers.NewExecutorController(executorService),
|
||||
File: controllers.NewFileController(constant.ScriptsWorkDir),
|
||||
Dashboard: controllers.NewDashboardController(cronService, executorService),
|
||||
Log: controllers.NewLogController(),
|
||||
Terminal: controllers.NewTerminalController(),
|
||||
Settings: controllers.NewSettingsController(userService),
|
||||
}
|
||||
}
|
||||
|
||||
// StopCron stops the cron service gracefully
|
||||
func StopCron() {
|
||||
if cronService != nil {
|
||||
cronService.Stop()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"io/fs"
|
||||
"net/http"
|
||||
|
||||
"baihu/internal/controllers"
|
||||
"baihu/internal/middleware"
|
||||
"baihu/internal/static"
|
||||
|
||||
"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
|
||||
Terminal *controllers.TerminalController
|
||||
Settings *controllers.SettingsController
|
||||
}
|
||||
|
||||
func mustSubFS(fsys fs.FS, dir string) fs.FS {
|
||||
sub, err := fs.Sub(fsys, dir)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return sub
|
||||
}
|
||||
|
||||
// cacheControl 返回设置 Cache-Control header 的中间件
|
||||
func cacheControl(value string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.Header("Cache-Control", value)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func Setup(c *Controllers) *gin.Engine {
|
||||
gin.SetMode(gin.ReleaseMode)
|
||||
router := gin.New()
|
||||
router.Use(middleware.GinLogger(), middleware.GinRecovery())
|
||||
|
||||
// Serve embedded Vue SPA static files with cache headers
|
||||
staticFS := static.GetFS()
|
||||
assetsGroup := router.Group("/assets")
|
||||
assetsGroup.Use(cacheControl("public, max-age=31536000, immutable")) // 1 year cache for hashed assets
|
||||
assetsGroup.StaticFS("/", http.FS(mustSubFS(staticFS, "assets")))
|
||||
|
||||
// Serve logo.svg with short cache
|
||||
router.GET("/logo.svg", func(ctx *gin.Context) {
|
||||
data, err := static.ReadFile("logo.svg")
|
||||
if err != nil {
|
||||
ctx.Status(404)
|
||||
return
|
||||
}
|
||||
ctx.Header("Cache-Control", "public, max-age=86400") // 1 day
|
||||
ctx.Data(200, "image/svg+xml", data)
|
||||
})
|
||||
|
||||
// SPA fallback - serve index.html (no cache for HTML)
|
||||
router.NoRoute(func(ctx *gin.Context) {
|
||||
data, err := static.ReadFile("index.html")
|
||||
if err != nil {
|
||||
ctx.String(500, "index.html not found")
|
||||
return
|
||||
}
|
||||
ctx.Header("Cache-Control", "no-cache, no-store, must-revalidate")
|
||||
ctx.Data(200, "text/html; charset=utf-8", data)
|
||||
})
|
||||
|
||||
// API routes
|
||||
api := router.Group("/api")
|
||||
{
|
||||
// Health check (无需认证)
|
||||
api.GET("/ping", func(ctx *gin.Context) {
|
||||
ctx.JSON(200, gin.H{"message": "pong"})
|
||||
})
|
||||
|
||||
// Authentication routes (无需认证)
|
||||
auth := api.Group("/auth")
|
||||
{
|
||||
auth.POST("/login", c.Auth.Login)
|
||||
auth.POST("/logout", c.Auth.Logout)
|
||||
auth.POST("/register", c.Auth.Register)
|
||||
}
|
||||
|
||||
// 需要认证的路由
|
||||
authorized := api.Group("")
|
||||
authorized.Use(middleware.AuthRequired())
|
||||
{
|
||||
// 获取当前用户
|
||||
authorized.GET("/auth/me", c.Auth.GetCurrentUser)
|
||||
|
||||
// Dashboard stats
|
||||
authorized.GET("/stats", c.Dashboard.GetStats)
|
||||
|
||||
// Task routes
|
||||
tasks := authorized.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)
|
||||
}
|
||||
|
||||
// Task execution routes
|
||||
execution := authorized.Group("/execute")
|
||||
{
|
||||
execution.POST("/task/:id", c.Executor.ExecuteTask)
|
||||
execution.POST("/command", c.Executor.ExecuteCommand)
|
||||
execution.GET("/results", c.Executor.GetLastResults)
|
||||
}
|
||||
|
||||
// Environment variable routes
|
||||
env := authorized.Group("/env")
|
||||
{
|
||||
env.POST("", c.Env.CreateEnvVar)
|
||||
env.GET("", c.Env.GetEnvVars)
|
||||
env.GET("/:id", c.Env.GetEnvVar)
|
||||
env.PUT("/:id", c.Env.UpdateEnvVar)
|
||||
env.DELETE("/:id", c.Env.DeleteEnvVar)
|
||||
}
|
||||
|
||||
// Script routes
|
||||
scripts := authorized.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)
|
||||
}
|
||||
|
||||
// File routes
|
||||
files := authorized.Group("/files")
|
||||
{
|
||||
files.GET("/tree", c.File.GetFileTree)
|
||||
files.GET("/content", c.File.GetFileContent)
|
||||
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("/upload", c.File.UploadArchive)
|
||||
files.POST("/upload-files", c.File.UploadFiles)
|
||||
}
|
||||
|
||||
// Log routes
|
||||
logs := authorized.Group("/logs")
|
||||
{
|
||||
logs.GET("", c.Log.GetLogs)
|
||||
logs.GET("/:id", c.Log.GetLogDetail)
|
||||
}
|
||||
|
||||
// Terminal routes
|
||||
authorized.GET("/terminal/ws", c.Terminal.HandleWebSocket)
|
||||
authorized.POST("/terminal/exec", c.Terminal.ExecuteShellCommand)
|
||||
|
||||
// Settings routes
|
||||
settings := authorized.Group("/settings")
|
||||
{
|
||||
settings.POST("/password", c.Settings.ChangePassword)
|
||||
settings.POST("/clean-logs", c.Settings.CleanLogs)
|
||||
settings.GET("/site", c.Settings.GetSiteSettings)
|
||||
settings.GET("/about", c.Settings.GetAbout)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return router
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"baihu/internal/constant"
|
||||
"encoding/json"
|
||||
"os"
|
||||
)
|
||||
|
||||
type ServerConfig struct {
|
||||
Port int `json:"port"`
|
||||
Host string `json:"host"`
|
||||
SiteName string `json:"site_name"`
|
||||
}
|
||||
|
||||
type DatabaseConfig struct {
|
||||
Type string `json:"type"`
|
||||
Host string `json:"host"`
|
||||
Port int `json:"port"`
|
||||
User string `json:"user"`
|
||||
Password string `json:"password"`
|
||||
DBName string `json:"dbname"`
|
||||
Path string `json:"path"`
|
||||
TablePrefix string `json:"table_prefix"`
|
||||
}
|
||||
|
||||
type SecurityConfig struct {
|
||||
JWTSecret string `json:"jwt_secret"`
|
||||
PasswordSalt string `json:"password_salt"`
|
||||
}
|
||||
|
||||
type TaskConfig struct {
|
||||
DefaultTimeout int `json:"default_timeout"`
|
||||
LogRetentionDays int `json:"log_retention_days"`
|
||||
}
|
||||
|
||||
type AppConfig struct {
|
||||
Server ServerConfig `json:"server"`
|
||||
Database DatabaseConfig `json:"database"`
|
||||
Security SecurityConfig `json:"security"`
|
||||
Task TaskConfig `json:"task"`
|
||||
}
|
||||
|
||||
var Config *AppConfig
|
||||
|
||||
func LoadConfig(path string) (*AppConfig, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
Config = &AppConfig{}
|
||||
if err := json.Unmarshal(data, Config); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 设置表前缀到 constant 包
|
||||
if Config.Database.TablePrefix != "" {
|
||||
constant.TablePrefix = Config.Database.TablePrefix
|
||||
}
|
||||
|
||||
// 设置 JWT 密钥
|
||||
if Config.Security.JWTSecret != "" {
|
||||
constant.JWTSecret = Config.Security.JWTSecret
|
||||
}
|
||||
|
||||
return Config, nil
|
||||
}
|
||||
|
||||
func GetConfig() *AppConfig {
|
||||
return Config
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/logger"
|
||||
"baihu/internal/models"
|
||||
|
||||
"github.com/robfig/cron/v3"
|
||||
)
|
||||
|
||||
// CronService manages scheduled tasks using robfig/cron
|
||||
type CronService struct {
|
||||
cron *cron.Cron
|
||||
taskService *TaskService
|
||||
executorService *ExecutorService
|
||||
entryMap map[uint]cron.EntryID // task ID -> cron entry ID
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewCronService creates a new cron service
|
||||
func NewCronService(taskService *TaskService, executorService *ExecutorService) *CronService {
|
||||
// 使用秒级精度的 cron parser,支持 5 位和 6 位表达式
|
||||
c := cron.New(cron.WithParser(cron.NewParser(
|
||||
cron.Minute | cron.Hour | cron.Dom | cron.Month | cron.Dow | cron.Descriptor,
|
||||
)))
|
||||
|
||||
return &CronService{
|
||||
cron: c,
|
||||
taskService: taskService,
|
||||
executorService: executorService,
|
||||
entryMap: make(map[uint]cron.EntryID),
|
||||
}
|
||||
}
|
||||
|
||||
// Start starts the cron service and loads all enabled tasks
|
||||
func (cs *CronService) Start() {
|
||||
cs.loadTasks()
|
||||
cs.cron.Start()
|
||||
logger.Info("Cron service started")
|
||||
}
|
||||
|
||||
// Stop stops the cron service
|
||||
func (cs *CronService) Stop() {
|
||||
ctx := cs.cron.Stop()
|
||||
<-ctx.Done()
|
||||
logger.Info("Cron service stopped")
|
||||
}
|
||||
|
||||
// loadTasks loads all enabled tasks from database
|
||||
func (cs *CronService) loadTasks() {
|
||||
tasks := cs.taskService.GetTasks()
|
||||
for _, task := range tasks {
|
||||
if task.Enabled {
|
||||
err := cs.AddTask(&task)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// AddTask adds a task to the cron scheduler
|
||||
func (cs *CronService) AddTask(task *models.Task) error {
|
||||
cs.mu.Lock()
|
||||
|
||||
// 如果已存在,先移除
|
||||
if entryID, exists := cs.entryMap[task.ID]; exists {
|
||||
cs.cron.Remove(entryID)
|
||||
delete(cs.entryMap, task.ID)
|
||||
}
|
||||
|
||||
taskID := task.ID
|
||||
entryID, err := cs.cron.AddFunc(task.Schedule, func() {
|
||||
cs.runTask(taskID)
|
||||
})
|
||||
if err != nil {
|
||||
cs.mu.Unlock()
|
||||
logger.Errorf("Failed to add task %d: %v", task.ID, err)
|
||||
return err
|
||||
}
|
||||
|
||||
cs.entryMap[task.ID] = entryID
|
||||
cs.mu.Unlock()
|
||||
|
||||
logger.Infof("Task %d (%s) scheduled with cron: %s", task.ID, task.Name, task.Schedule)
|
||||
|
||||
// 更新下次运行时间
|
||||
cs.updateNextRun(task.ID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveTask removes a task from the cron scheduler
|
||||
func (cs *CronService) RemoveTask(taskID uint) {
|
||||
cs.mu.Lock()
|
||||
defer cs.mu.Unlock()
|
||||
|
||||
if entryID, exists := cs.entryMap[taskID]; exists {
|
||||
cs.cron.Remove(entryID)
|
||||
delete(cs.entryMap, taskID)
|
||||
logger.Infof("Task %d removed from scheduler", taskID)
|
||||
}
|
||||
}
|
||||
|
||||
// runTask executes a task and updates its status
|
||||
func (cs *CronService) runTask(taskID uint) {
|
||||
logger.Infof("Running task %d", taskID)
|
||||
|
||||
// 更新 last_run
|
||||
now := time.Now()
|
||||
database.DB.Model(&models.Task{}).Where("id = ?", taskID).Update("last_run", now)
|
||||
|
||||
// 执行任务
|
||||
cs.executorService.ExecuteTask(int(taskID))
|
||||
|
||||
// 更新 next_run
|
||||
cs.updateNextRun(taskID)
|
||||
}
|
||||
|
||||
// updateNextRun updates the next run time for a task
|
||||
func (cs *CronService) updateNextRun(taskID uint) {
|
||||
cs.mu.RLock()
|
||||
entryID, exists := cs.entryMap[taskID]
|
||||
cs.mu.RUnlock()
|
||||
|
||||
if !exists {
|
||||
return
|
||||
}
|
||||
|
||||
entry := cs.cron.Entry(entryID)
|
||||
if !entry.Next.IsZero() {
|
||||
database.DB.Model(&models.Task{}).Where("id = ?", taskID).Update("next_run", entry.Next)
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateCron validates a cron expression
|
||||
func (cs *CronService) ValidateCron(expression string) error {
|
||||
parser := cron.NewParser(cron.Minute | cron.Hour | cron.Dom | cron.Month | cron.Dow | cron.Descriptor)
|
||||
_, err := parser.Parse(expression)
|
||||
return err
|
||||
}
|
||||
|
||||
// GetScheduledCount returns the number of scheduled tasks
|
||||
func (cs *CronService) GetScheduledCount() int {
|
||||
cs.mu.RLock()
|
||||
defer cs.mu.RUnlock()
|
||||
return len(cs.entryMap)
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/models"
|
||||
)
|
||||
|
||||
type EnvService struct{}
|
||||
|
||||
func NewEnvService() *EnvService {
|
||||
return &EnvService{}
|
||||
}
|
||||
|
||||
func (es *EnvService) CreateEnvVar(name, value, remark string, userID int) *models.EnvironmentVariable {
|
||||
env := &models.EnvironmentVariable{
|
||||
Name: name,
|
||||
Value: value,
|
||||
Remark: remark,
|
||||
UserID: uint(userID),
|
||||
}
|
||||
database.DB.Create(env)
|
||||
return env
|
||||
}
|
||||
|
||||
func (es *EnvService) GetEnvVarsByUserID(userID int) []models.EnvironmentVariable {
|
||||
var envs []models.EnvironmentVariable
|
||||
database.DB.Where("user_id = ?", userID).Find(&envs)
|
||||
return envs
|
||||
}
|
||||
|
||||
func (es *EnvService) GetEnvVarsWithPagination(userID int, name string, page, pageSize int) ([]models.EnvironmentVariable, int64) {
|
||||
var envs []models.EnvironmentVariable
|
||||
var total int64
|
||||
|
||||
query := database.DB.Model(&models.EnvironmentVariable{}).Where("user_id = ?", userID)
|
||||
if name != "" {
|
||||
query = query.Where("name LIKE ?", "%"+name+"%")
|
||||
}
|
||||
|
||||
query.Count(&total)
|
||||
query.Order("id DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&envs)
|
||||
return envs, total
|
||||
}
|
||||
|
||||
func (es *EnvService) GetEnvVarByID(id int) *models.EnvironmentVariable {
|
||||
var env models.EnvironmentVariable
|
||||
if err := database.DB.First(&env, id).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return &env
|
||||
}
|
||||
|
||||
func (es *EnvService) UpdateEnvVar(id int, name, value, remark string) *models.EnvironmentVariable {
|
||||
var env models.EnvironmentVariable
|
||||
if err := database.DB.First(&env, id).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
env.Name = name
|
||||
env.Value = value
|
||||
env.Remark = remark
|
||||
database.DB.Save(&env)
|
||||
return &env
|
||||
}
|
||||
|
||||
func (es *EnvService) DeleteEnvVar(id int) bool {
|
||||
result := database.DB.Delete(&models.EnvironmentVariable{}, id)
|
||||
return result.RowsAffected > 0
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"os/exec"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"baihu/internal/constant"
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/logger"
|
||||
"baihu/internal/models"
|
||||
"baihu/internal/utils"
|
||||
)
|
||||
|
||||
// ExecutionResult represents the result of a task execution
|
||||
type ExecutionResult struct {
|
||||
TaskID int
|
||||
Success bool
|
||||
Output string
|
||||
Error string
|
||||
Start time.Time
|
||||
End time.Time
|
||||
}
|
||||
|
||||
// ExecutorService handles task execution
|
||||
type ExecutorService struct {
|
||||
taskService *TaskService
|
||||
results []ExecutionResult
|
||||
runningTasks map[int]bool // 正在运行的任务
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewExecutorService creates a new executor service
|
||||
func NewExecutorService(taskService *TaskService) *ExecutorService {
|
||||
return &ExecutorService{
|
||||
taskService: taskService,
|
||||
results: make([]ExecutionResult, 0),
|
||||
runningTasks: make(map[int]bool),
|
||||
}
|
||||
}
|
||||
|
||||
// ExecuteTask executes a task by ID
|
||||
func (es *ExecutorService) ExecuteTask(taskID int) *ExecutionResult {
|
||||
task := es.taskService.GetTaskByID(taskID)
|
||||
if task == nil {
|
||||
return &ExecutionResult{
|
||||
TaskID: taskID,
|
||||
Success: false,
|
||||
Error: "Task not found",
|
||||
Start: time.Now(),
|
||||
End: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
// 标记任务开始运行
|
||||
es.mu.Lock()
|
||||
es.runningTasks[taskID] = true
|
||||
es.mu.Unlock()
|
||||
|
||||
// 使用任务配置的超时时间
|
||||
timeout := task.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = constant.DefaultTaskTimeout
|
||||
}
|
||||
result := es.ExecuteCommandWithTimeout(task.Command, time.Duration(timeout)*time.Minute)
|
||||
result.TaskID = taskID
|
||||
|
||||
// 标记任务结束
|
||||
es.mu.Lock()
|
||||
delete(es.runningTasks, taskID)
|
||||
es.mu.Unlock()
|
||||
|
||||
// Save log to database
|
||||
es.saveTaskLog(uint(taskID), task.Command, result)
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetRunningCount 获取正在运行的任务数量
|
||||
func (es *ExecutorService) GetRunningCount() int {
|
||||
es.mu.RLock()
|
||||
defer es.mu.RUnlock()
|
||||
return len(es.runningTasks)
|
||||
}
|
||||
|
||||
// ExecuteCommand executes a shell command with default timeout
|
||||
func (es *ExecutorService) ExecuteCommand(command string) *ExecutionResult {
|
||||
return es.ExecuteCommandWithTimeout(command, time.Duration(constant.DefaultTaskTimeout)*time.Minute)
|
||||
}
|
||||
|
||||
// ExecuteCommandWithTimeout executes a shell command with specified timeout
|
||||
func (es *ExecutorService) ExecuteCommandWithTimeout(command string, timeout time.Duration) *ExecutionResult {
|
||||
result := &ExecutionResult{
|
||||
Success: false,
|
||||
Start: time.Now(),
|
||||
}
|
||||
|
||||
// Create a context with timeout
|
||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
||||
defer cancel()
|
||||
|
||||
// Execute the command
|
||||
shell, args := utils.GetShellCommand(command)
|
||||
cmd := exec.CommandContext(ctx, shell, args...)
|
||||
var stdout, stderr bytes.Buffer
|
||||
cmd.Stdout = &stdout
|
||||
cmd.Stderr = &stderr
|
||||
|
||||
err := cmd.Run()
|
||||
result.End = time.Now()
|
||||
|
||||
// Process results
|
||||
result.Output = stdout.String()
|
||||
if err != nil {
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
result.Error = "执行超时\n" + stderr.String()
|
||||
} else {
|
||||
result.Error = err.Error() + "\n" + stderr.String()
|
||||
}
|
||||
} else {
|
||||
result.Success = true
|
||||
}
|
||||
|
||||
// Store result
|
||||
es.mu.Lock()
|
||||
es.results = append(es.results, *result)
|
||||
// Keep only the last 100 results to prevent memory issues
|
||||
if len(es.results) > 100 {
|
||||
es.results = es.results[1:]
|
||||
}
|
||||
es.mu.Unlock()
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// GetLastResults returns the last execution results
|
||||
func (es *ExecutorService) GetLastResults(count int) []ExecutionResult {
|
||||
es.mu.RLock()
|
||||
defer es.mu.RUnlock()
|
||||
|
||||
start := 0
|
||||
if len(es.results) > count {
|
||||
start = len(es.results) - count
|
||||
}
|
||||
|
||||
results := make([]ExecutionResult, len(es.results[start:]))
|
||||
copy(results, es.results[start:])
|
||||
return results
|
||||
}
|
||||
|
||||
// saveTaskLog saves execution log to database with gzip+base64 compression
|
||||
func (es *ExecutorService) saveTaskLog(taskID uint, command string, result *ExecutionResult) {
|
||||
output := result.Output
|
||||
if result.Error != "" {
|
||||
output += "\n[ERROR]\n" + result.Error
|
||||
}
|
||||
|
||||
// Compress output
|
||||
compressed, err := utils.CompressToBase64(output)
|
||||
if err != nil {
|
||||
logger.Errorf("Failed to compress log: %v", err)
|
||||
compressed = ""
|
||||
}
|
||||
|
||||
status := "success"
|
||||
if !result.Success {
|
||||
status = "failed"
|
||||
}
|
||||
|
||||
taskLog := &models.TaskLog{
|
||||
TaskID: taskID,
|
||||
Command: command,
|
||||
Output: compressed,
|
||||
Status: status,
|
||||
Duration: result.End.Sub(result.Start).Milliseconds(),
|
||||
}
|
||||
|
||||
if err := database.DB.Create(taskLog).Error; err != nil {
|
||||
logger.Errorf("Failed to save task log: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"baihu/internal/logger"
|
||||
)
|
||||
|
||||
const (
|
||||
InitSection = "system"
|
||||
InitKey = "initialized"
|
||||
InitValue = "true"
|
||||
)
|
||||
|
||||
type InitService struct {
|
||||
settingsService *SettingsService
|
||||
userService *UserService
|
||||
}
|
||||
|
||||
func NewInitService(settingsService *SettingsService, userService *UserService) *InitService {
|
||||
return &InitService{
|
||||
settingsService: settingsService,
|
||||
userService: userService,
|
||||
}
|
||||
}
|
||||
|
||||
// Initialize 执行初始化,如果已初始化则跳过
|
||||
func (s *InitService) Initialize() {
|
||||
if s.IsInitialized() {
|
||||
logger.Info("系统已初始化,跳过")
|
||||
return
|
||||
}
|
||||
|
||||
logger.Info("开始初始化系统...")
|
||||
|
||||
// 创建管理员账号
|
||||
s.createAdminUser()
|
||||
|
||||
// 标记为已初始化
|
||||
s.settingsService.Set(InitSection, InitKey, InitValue)
|
||||
logger.Info("系统初始化完成")
|
||||
}
|
||||
|
||||
// IsInitialized 检查是否已初始化
|
||||
func (s *InitService) IsInitialized() bool {
|
||||
return s.settingsService.Get(InitSection, InitKey) == InitValue
|
||||
}
|
||||
|
||||
// createAdminUser 创建管理员账号
|
||||
func (s *InitService) createAdminUser() {
|
||||
existingUser := s.userService.GetUserByUsername("admin")
|
||||
if existingUser != nil {
|
||||
logger.Info("管理员账号已存在,跳过创建")
|
||||
return
|
||||
}
|
||||
|
||||
s.userService.CreateUser("admin", "123456", "admin@local", "admin")
|
||||
logger.Info("管理员账号创建成功: admin / 123456")
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/models"
|
||||
)
|
||||
|
||||
type ScriptService struct{}
|
||||
|
||||
func NewScriptService() *ScriptService {
|
||||
return &ScriptService{}
|
||||
}
|
||||
|
||||
func (ss *ScriptService) CreateScript(name, content string, userID int) *models.Script {
|
||||
script := &models.Script{
|
||||
Name: name,
|
||||
Content: content,
|
||||
UserID: uint(userID),
|
||||
}
|
||||
database.DB.Create(script)
|
||||
return script
|
||||
}
|
||||
|
||||
func (ss *ScriptService) GetScriptsByUserID(userID int) []models.Script {
|
||||
var scripts []models.Script
|
||||
database.DB.Where("user_id = ?", userID).Find(&scripts)
|
||||
return scripts
|
||||
}
|
||||
|
||||
func (ss *ScriptService) GetScriptByID(id int) *models.Script {
|
||||
var script models.Script
|
||||
if err := database.DB.First(&script, id).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return &script
|
||||
}
|
||||
|
||||
func (ss *ScriptService) UpdateScript(id int, name, content string) *models.Script {
|
||||
var script models.Script
|
||||
if err := database.DB.First(&script, id).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
script.Name = name
|
||||
script.Content = content
|
||||
database.DB.Save(&script)
|
||||
return &script
|
||||
}
|
||||
|
||||
func (ss *ScriptService) DeleteScript(id int) bool {
|
||||
result := database.DB.Delete(&models.Script{}, id)
|
||||
return result.RowsAffected > 0
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/models"
|
||||
)
|
||||
|
||||
type SettingsService struct{}
|
||||
|
||||
func NewSettingsService() *SettingsService {
|
||||
return &SettingsService{}
|
||||
}
|
||||
|
||||
// Get 获取设置值
|
||||
func (s *SettingsService) Get(section, key string) string {
|
||||
var setting models.Setting
|
||||
if err := database.DB.Where("section = ? AND key = ?", section, key).First(&setting).Error; err != nil {
|
||||
return ""
|
||||
}
|
||||
return setting.Value
|
||||
}
|
||||
|
||||
// GetWithDefault 获取设置值,如果不存在则返回默认值
|
||||
func (s *SettingsService) GetWithDefault(section, key, defaultValue string) string {
|
||||
value := s.Get(section, key)
|
||||
if value == "" {
|
||||
return defaultValue
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
// Set 设置值
|
||||
func (s *SettingsService) Set(section, key, value string) error {
|
||||
var setting models.Setting
|
||||
result := database.DB.Where("section = ? AND key = ?", section, key).First(&setting)
|
||||
if result.Error != nil {
|
||||
// 不存在则创建
|
||||
setting = models.Setting{
|
||||
Section: section,
|
||||
Key: key,
|
||||
Value: value,
|
||||
}
|
||||
return database.DB.Create(&setting).Error
|
||||
}
|
||||
// 存在则更新
|
||||
return database.DB.Model(&setting).Update("value", value).Error
|
||||
}
|
||||
|
||||
// GetBySection 获取某个 section 下的所有设置
|
||||
func (s *SettingsService) GetBySection(section string) map[string]string {
|
||||
var settings []models.Setting
|
||||
database.DB.Where("section = ?", section).Find(&settings)
|
||||
|
||||
result := make(map[string]string)
|
||||
for _, s := range settings {
|
||||
result[s.Key] = s.Value
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// Delete 删除设置
|
||||
func (s *SettingsService) Delete(section, key string) error {
|
||||
return database.DB.Where("section = ? AND key = ?", section, key).Delete(&models.Setting{}).Error
|
||||
}
|
||||
|
||||
// DeleteBySection 删除某个 section 下的所有设置
|
||||
func (s *SettingsService) DeleteBySection(section string) error {
|
||||
return database.DB.Where("section = ?", section).Delete(&models.Setting{}).Error
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/models"
|
||||
)
|
||||
|
||||
type TaskService struct{}
|
||||
|
||||
func NewTaskService() *TaskService {
|
||||
return &TaskService{}
|
||||
}
|
||||
|
||||
func (ts *TaskService) CreateTask(name, command, schedule string) *models.Task {
|
||||
task := &models.Task{
|
||||
Name: name,
|
||||
Command: command,
|
||||
Schedule: schedule,
|
||||
Enabled: true,
|
||||
}
|
||||
database.DB.Create(task)
|
||||
return task
|
||||
}
|
||||
|
||||
func (ts *TaskService) GetTasks() []models.Task {
|
||||
var tasks []models.Task
|
||||
database.DB.Find(&tasks)
|
||||
return tasks
|
||||
}
|
||||
|
||||
// GetTasksWithPagination 分页获取任务列表
|
||||
func (ts *TaskService) GetTasksWithPagination(page, pageSize int, name string) ([]models.Task, int64) {
|
||||
var tasks []models.Task
|
||||
var total int64
|
||||
|
||||
query := database.DB.Model(&models.Task{})
|
||||
if name != "" {
|
||||
query = query.Where("name LIKE ?", "%"+name+"%")
|
||||
}
|
||||
|
||||
query.Count(&total)
|
||||
query.Order("id DESC").Offset((page - 1) * pageSize).Limit(pageSize).Find(&tasks)
|
||||
|
||||
return tasks, total
|
||||
}
|
||||
|
||||
func (ts *TaskService) GetTaskByID(id int) *models.Task {
|
||||
var task models.Task
|
||||
if err := database.DB.First(&task, id).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return &task
|
||||
}
|
||||
|
||||
func (ts *TaskService) UpdateTask(id int, name, command, schedule string, enabled bool) *models.Task {
|
||||
var task models.Task
|
||||
if err := database.DB.First(&task, id).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
task.Name = name
|
||||
task.Command = command
|
||||
task.Schedule = schedule
|
||||
task.Enabled = enabled
|
||||
database.DB.Save(&task)
|
||||
return &task
|
||||
}
|
||||
|
||||
func (ts *TaskService) DeleteTask(id int) bool {
|
||||
result := database.DB.Delete(&models.Task{}, id)
|
||||
return result.RowsAffected > 0
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
|
||||
"baihu/internal/database"
|
||||
"baihu/internal/models"
|
||||
)
|
||||
|
||||
type UserService struct{}
|
||||
|
||||
func NewUserService() *UserService {
|
||||
return &UserService{}
|
||||
}
|
||||
|
||||
func (us *UserService) hashPassword(password string) string {
|
||||
salt := ""
|
||||
if Config != nil {
|
||||
salt = Config.Security.PasswordSalt
|
||||
}
|
||||
hash := sha256.Sum256([]byte(password + salt))
|
||||
return hex.EncodeToString(hash[:])
|
||||
}
|
||||
|
||||
func (us *UserService) CreateUser(username, password, email, role string) *models.User {
|
||||
user := &models.User{
|
||||
Username: username,
|
||||
Password: us.hashPassword(password),
|
||||
Email: email,
|
||||
Role: role,
|
||||
}
|
||||
database.DB.Create(user)
|
||||
return user
|
||||
}
|
||||
|
||||
func (us *UserService) GetUserByUsername(username string) *models.User {
|
||||
var user models.User
|
||||
if err := database.DB.Where("username = ?", username).First(&user).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return &user
|
||||
}
|
||||
|
||||
func (us *UserService) ValidatePassword(user *models.User, password string) bool {
|
||||
return user.Password == us.hashPassword(password)
|
||||
}
|
||||
|
||||
func (us *UserService) EnsureAdminExists() {
|
||||
var count int64
|
||||
database.DB.Model(&models.User{}).Where("role = ?", "admin").Count(&count)
|
||||
if count == 0 {
|
||||
us.CreateUser("admin", "admin123", "admin@local", "admin")
|
||||
}
|
||||
}
|
||||
|
||||
func (us *UserService) AuthenticateUser(username, password string) bool {
|
||||
user := us.GetUserByUsername(username)
|
||||
if user == nil {
|
||||
return false
|
||||
}
|
||||
return us.ValidatePassword(user, password)
|
||||
}
|
||||
|
||||
func (us *UserService) UpdatePassword(userID uint, newPassword string) error {
|
||||
return database.DB.Model(&models.User{}).Where("id = ?", userID).Update("password", us.hashPassword(newPassword)).Error
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
package static
|
||||
|
||||
import (
|
||||
"embed"
|
||||
"io/fs"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
//go:embed dist/*
|
||||
var distFS embed.FS
|
||||
|
||||
// GetFileSystem 返回嵌入的静态文件系统
|
||||
func GetFileSystem() http.FileSystem {
|
||||
subFS, err := fs.Sub(distFS, "dist")
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return http.FS(subFS)
|
||||
}
|
||||
|
||||
// GetFS 返回嵌入的 fs.FS
|
||||
func GetFS() fs.FS {
|
||||
subFS, err := fs.Sub(distFS, "dist")
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return subFS
|
||||
}
|
||||
|
||||
// ReadFile 读取嵌入的文件
|
||||
func ReadFile(name string) ([]byte, error) {
|
||||
return distFS.ReadFile("dist/" + name)
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"archive/zip"
|
||||
"compress/gzip"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func ExtractZip(src, dest string) error {
|
||||
r, err := zip.OpenReader(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
for _, f := range r.File {
|
||||
fpath := filepath.Join(dest, f.Name)
|
||||
|
||||
// 安全检查:防止路径遍历
|
||||
if !strings.HasPrefix(fpath, filepath.Clean(dest)+string(os.PathSeparator)) {
|
||||
continue
|
||||
}
|
||||
|
||||
if f.FileInfo().IsDir() {
|
||||
os.MkdirAll(fpath, 0755)
|
||||
continue
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(fpath), 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
outFile, err := os.OpenFile(fpath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, f.Mode())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
rc, err := f.Open()
|
||||
if err != nil {
|
||||
outFile.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
_, err = io.Copy(outFile, rc)
|
||||
outFile.Close()
|
||||
rc.Close()
|
||||
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ExtractTar(src, dest string) error {
|
||||
file, err := os.Open(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
return extractTarReader(tar.NewReader(file), dest)
|
||||
}
|
||||
|
||||
func ExtractTarGz(src, dest string) error {
|
||||
file, err := os.Open(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
gzr, err := gzip.NewReader(file)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer gzr.Close()
|
||||
|
||||
return extractTarReader(tar.NewReader(gzr), dest)
|
||||
}
|
||||
|
||||
func extractTarReader(tr *tar.Reader, dest string) error {
|
||||
for {
|
||||
header, err := tr.Next()
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
fpath := filepath.Join(dest, header.Name)
|
||||
|
||||
// 安全检查:防止路径遍历
|
||||
if !strings.HasPrefix(fpath, filepath.Clean(dest)+string(os.PathSeparator)) {
|
||||
continue
|
||||
}
|
||||
|
||||
switch header.Typeflag {
|
||||
case tar.TypeDir:
|
||||
os.MkdirAll(fpath, 0755)
|
||||
case tar.TypeReg:
|
||||
if err := os.MkdirAll(filepath.Dir(fpath), 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
outFile, err := os.Create(fpath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := io.Copy(outFile, tr); err != nil {
|
||||
outFile.Close()
|
||||
return err
|
||||
}
|
||||
outFile.Close()
|
||||
|
||||
os.Chmod(fpath, os.FileMode(header.Mode))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"compress/gzip"
|
||||
"encoding/base64"
|
||||
"io"
|
||||
)
|
||||
|
||||
// CompressToBase64 compresses data using gzip and encodes to base64
|
||||
func CompressToBase64(data string) (string, error) {
|
||||
var buf bytes.Buffer
|
||||
gz := gzip.NewWriter(&buf)
|
||||
if _, err := gz.Write([]byte(data)); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := gz.Close(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.StdEncoding.EncodeToString(buf.Bytes()), nil
|
||||
}
|
||||
|
||||
// DecompressFromBase64 decodes base64 and decompresses gzip data
|
||||
func DecompressFromBase64(data string) (string, error) {
|
||||
decoded, err := base64.StdEncoding.DecodeString(data)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
gz, err := gzip.NewReader(bytes.NewReader(decoded))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer gz.Close()
|
||||
result, err := io.ReadAll(gz)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(result), nil
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
|
||||
"baihu/internal/constant"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// Pagination 分页参数
|
||||
type Pagination struct {
|
||||
Page int
|
||||
PageSize int
|
||||
}
|
||||
|
||||
// ParsePagination 从请求中解析分页参数
|
||||
func ParsePagination(c *gin.Context) Pagination {
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", strconv.Itoa(constant.DefaultPageSize)))
|
||||
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 || pageSize > 100 {
|
||||
pageSize = constant.DefaultPageSize
|
||||
}
|
||||
|
||||
return Pagination{Page: page, PageSize: pageSize}
|
||||
}
|
||||
|
||||
// Offset 计算偏移量
|
||||
func (p Pagination) Offset() int {
|
||||
return (p.Page - 1) * p.PageSize
|
||||
}
|
||||
|
||||
// PaginatedResponse 分页响应
|
||||
func PaginatedResponse(c *gin.Context, data interface{}, total int64, p Pagination) {
|
||||
Success(c, gin.H{
|
||||
"data": data,
|
||||
"total": total,
|
||||
"page": p.Page,
|
||||
"page_size": p.PageSize,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type Response struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
Data interface{} `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
func Success(c *gin.Context, data interface{}) {
|
||||
c.JSON(http.StatusOK, Response{
|
||||
Code: 200,
|
||||
Msg: "success",
|
||||
Data: data,
|
||||
})
|
||||
}
|
||||
|
||||
func SuccessMsg(c *gin.Context, msg string) {
|
||||
c.JSON(http.StatusOK, Response{
|
||||
Code: 200,
|
||||
Msg: msg,
|
||||
})
|
||||
}
|
||||
|
||||
func Error(c *gin.Context, code int, msg string) {
|
||||
c.JSON(http.StatusOK, Response{
|
||||
Code: code,
|
||||
Msg: msg,
|
||||
})
|
||||
}
|
||||
|
||||
func BadRequest(c *gin.Context, msg string) {
|
||||
Error(c, 400, msg)
|
||||
}
|
||||
|
||||
func Unauthorized(c *gin.Context, msg string) {
|
||||
Error(c, 401, msg)
|
||||
}
|
||||
|
||||
func Forbidden(c *gin.Context, msg string) {
|
||||
Error(c, 403, msg)
|
||||
}
|
||||
|
||||
func NotFound(c *gin.Context, msg string) {
|
||||
Error(c, 404, msg)
|
||||
}
|
||||
|
||||
func ServerError(c *gin.Context, msg string) {
|
||||
Error(c, 500, msg)
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
)
|
||||
|
||||
// GetShell 返回当前操作系统的 shell 和参数
|
||||
func GetShell() (shell string, args []string) {
|
||||
if runtime.GOOS == "windows" {
|
||||
return "cmd", []string{}
|
||||
}
|
||||
|
||||
// 优先使用环境变量中的 SHELL
|
||||
if envShell := os.Getenv("SHELL"); envShell != "" {
|
||||
return envShell, []string{}
|
||||
}
|
||||
|
||||
// macOS 默认使用 zsh
|
||||
if runtime.GOOS == "darwin" {
|
||||
if _, err := exec.LookPath("/bin/zsh"); err == nil {
|
||||
return "/bin/zsh", []string{}
|
||||
}
|
||||
}
|
||||
|
||||
// Linux 默认使用 bash
|
||||
return "/bin/bash", []string{}
|
||||
}
|
||||
|
||||
// GetShellCommand 返回执行命令的 shell 和参数
|
||||
func GetShellCommand(command string) (shell string, args []string) {
|
||||
shell, _ = GetShell()
|
||||
if runtime.GOOS == "windows" {
|
||||
return shell, []string{"/c", command}
|
||||
}
|
||||
return shell, []string{"-c", command}
|
||||
}
|
||||
|
||||
// NewShellCmd 创建一个交互式 shell 命令
|
||||
func NewShellCmd() *exec.Cmd {
|
||||
shell, _ := GetShell()
|
||||
if runtime.GOOS == "windows" {
|
||||
return exec.Command(shell)
|
||||
}
|
||||
// Unix 系统使用 -i 启用交互模式
|
||||
return exec.Command(shell, "-i")
|
||||
}
|
||||
|
||||
// NewShellCommandCmd 创建一个执行指定命令的 shell 命令
|
||||
func NewShellCommandCmd(command string) *exec.Cmd {
|
||||
shell, args := GetShellCommand(command)
|
||||
return exec.Command(shell, args...)
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"baihu/internal/constant"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
type Claims struct {
|
||||
UserID uint `json:"user_id"`
|
||||
Username string `json:"username"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// GenerateToken 生成 JWT token
|
||||
func GenerateToken(userID uint, username string) (string, error) {
|
||||
claims := Claims{
|
||||
UserID: userID,
|
||||
Username: username,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Duration(constant.TokenExpireDays) * 24 * time.Hour)),
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
},
|
||||
}
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return token.SignedString([]byte(constant.JWTSecret))
|
||||
}
|
||||
|
||||
// ParseToken 解析 JWT token
|
||||
func ParseToken(tokenString string) (uint, string, error) {
|
||||
token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (any, error) {
|
||||
return []byte(constant.JWTSecret), nil
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return 0, "", err
|
||||
}
|
||||
|
||||
if claims, ok := token.Claims.(*Claims); ok && token.Valid {
|
||||
return claims.UserID, claims.Username, nil
|
||||
}
|
||||
|
||||
return 0, "", errors.New("invalid token")
|
||||
}
|
||||
Reference in New Issue
Block a user