feat: refact id column define
This commit is contained in:
@@ -14,6 +14,7 @@ import (
|
||||
"github.com/engigu/baihu-panel/internal/logger"
|
||||
"github.com/engigu/baihu-panel/internal/models"
|
||||
"github.com/engigu/baihu-panel/internal/services/tasks"
|
||||
"github.com/engigu/baihu-panel/internal/utils"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
@@ -46,6 +47,7 @@ func (s *AgentService) CreateToken(remark string, maxUses int, expiresAt *time.T
|
||||
token := generateToken()
|
||||
|
||||
agentToken := &models.AgentToken{
|
||||
ID: utils.GenerateID(),
|
||||
Token: token,
|
||||
Remark: remark,
|
||||
MaxUses: maxUses,
|
||||
@@ -69,8 +71,8 @@ func (s *AgentService) ListTokens() []models.AgentToken {
|
||||
}
|
||||
|
||||
// DeleteToken 删除令牌
|
||||
func (s *AgentService) DeleteToken(id uint) error {
|
||||
return database.DB.Delete(&models.AgentToken{}, id).Error
|
||||
func (s *AgentService) DeleteToken(id string) error {
|
||||
return database.DB.Where("id = ?", id).Delete(&models.AgentToken{}).Error
|
||||
}
|
||||
|
||||
// ValidateToken 验证令牌
|
||||
@@ -98,7 +100,7 @@ func (s *AgentService) ValidateToken(token string) (*models.AgentToken, error) {
|
||||
}
|
||||
|
||||
// UseToken 使用令牌(增加使用计数)
|
||||
func (s *AgentService) UseToken(id uint) {
|
||||
func (s *AgentService) UseToken(id string) {
|
||||
database.DB.Model(&models.AgentToken{}).Where("id = ?", id).UpdateColumn("used_count", gorm.Expr("used_count + 1"))
|
||||
}
|
||||
|
||||
@@ -126,7 +128,7 @@ func (s *AgentService) RegisterByToken(token string, machineID string, ip string
|
||||
"last_seen": now,
|
||||
})
|
||||
s.UseToken(agentToken.ID)
|
||||
logger.Infof("[Agent] Agent #%d 通过 machine_id 复用 (%s)", existing.ID, machineID[:8]+"...")
|
||||
logger.Infof("[Agent] Agent #%s 通过 machine_id 复用 (%s)", existing.ID, machineID[:8]+"...")
|
||||
return &existing, false, nil
|
||||
}
|
||||
}
|
||||
@@ -134,6 +136,7 @@ func (s *AgentService) RegisterByToken(token string, machineID string, ip string
|
||||
// 创建 Agent,使用令牌作为认证 Token
|
||||
now := models.LocalTime(time.Now())
|
||||
agent := &models.Agent{
|
||||
ID: utils.GenerateID(),
|
||||
Name: fmt.Sprintf("agent-%d", time.Now().Unix()),
|
||||
Token: token,
|
||||
MachineID: machineID,
|
||||
@@ -148,7 +151,7 @@ func (s *AgentService) RegisterByToken(token string, machineID string, ip string
|
||||
}
|
||||
|
||||
s.UseToken(agentToken.ID)
|
||||
logger.Infof("[Agent] Agent 通过令牌注册: #%d (%s)", agent.ID, ip)
|
||||
logger.Infof("[Agent] Agent 通过令牌注册: #%s (%s)", agent.ID, ip)
|
||||
return agent, true, nil
|
||||
}
|
||||
|
||||
@@ -173,6 +176,7 @@ func (s *AgentService) Register(req *models.AgentRegisterRequest, ip string) (*m
|
||||
// 创建新 Agent,使用令牌作为认证 Token
|
||||
now := models.LocalTime(time.Now())
|
||||
agent := &models.Agent{
|
||||
ID: utils.GenerateID(),
|
||||
Name: req.Name,
|
||||
Token: req.Token,
|
||||
Hostname: req.Hostname,
|
||||
@@ -194,7 +198,7 @@ func (s *AgentService) Register(req *models.AgentRegisterRequest, ip string) (*m
|
||||
}
|
||||
|
||||
// Update 更新 Agent
|
||||
func (s *AgentService) Update(id uint, name, description string, enabled bool) error {
|
||||
func (s *AgentService) Update(id string, name, description string, enabled bool) error {
|
||||
return database.DB.Model(&models.Agent{}).Where("id = ?", id).Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"description": description,
|
||||
@@ -203,7 +207,7 @@ func (s *AgentService) Update(id uint, name, description string, enabled bool) e
|
||||
}
|
||||
|
||||
// Delete 删除 Agent(物理删除)
|
||||
func (s *AgentService) Delete(id uint) error {
|
||||
func (s *AgentService) Delete(id string) error {
|
||||
// 检查是否有关联任务
|
||||
var count int64
|
||||
database.DB.Model(&models.Task{}).Where("agent_id = ?", id).Count(&count)
|
||||
@@ -211,13 +215,13 @@ func (s *AgentService) Delete(id uint) error {
|
||||
return &ServiceError{Message: "该 Agent 下还有关联任务,无法删除"}
|
||||
}
|
||||
|
||||
return database.DB.Unscoped().Delete(&models.Agent{}, id).Error
|
||||
return database.DB.Unscoped().Where("id = ?", id).Delete(&models.Agent{}).Error
|
||||
}
|
||||
|
||||
// GetByID 根据 ID 获取 Agent
|
||||
func (s *AgentService) GetByID(id uint) *models.Agent {
|
||||
func (s *AgentService) GetByID(id string) *models.Agent {
|
||||
var agent models.Agent
|
||||
if err := database.DB.First(&agent, id).Error; err != nil {
|
||||
if err := database.DB.Where("id = ?", id).First(&agent).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return &agent
|
||||
@@ -249,7 +253,7 @@ func (s *AgentService) List() []models.Agent {
|
||||
}
|
||||
|
||||
// RegenerateToken 重新生成 Token - 已废弃,保留空实现避免路由错误
|
||||
func (s *AgentService) RegenerateToken(id uint) (string, error) {
|
||||
func (s *AgentService) RegenerateToken(id string) (string, error) {
|
||||
return "", &ServiceError{Message: "此功能已禁用"}
|
||||
}
|
||||
|
||||
@@ -301,7 +305,7 @@ func (s *AgentService) Heartbeat(token, ip, version, buildTime, hostname, osType
|
||||
}
|
||||
|
||||
// GetTasks 获取 Agent 的任务列表
|
||||
func (s *AgentService) GetTasks(agentID uint) []models.AgentTask {
|
||||
func (s *AgentService) GetTasks(agentID string) []models.AgentTask {
|
||||
var tasks []models.Task
|
||||
database.DB.Where("agent_id = ? AND enabled = ?", agentID, true).Find(&tasks)
|
||||
|
||||
@@ -359,13 +363,13 @@ func (s *AgentService) ReportResult(result *models.AgentTaskResult) error {
|
||||
|
||||
// 先尝试通知正在等待的 goroutine
|
||||
if agentWSManager.NotifyRemoteResult(result) {
|
||||
logger.Infof("[Agent] 已通知正在等待任务 #%d 结果的 goroutine", result.TaskID)
|
||||
logger.Infof("[Agent] 已通知正在等待任务 #%s 结果的 goroutine", result.TaskID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// 如果没有人在等待(例如服务重启后),则由本协程负责处理结果入库
|
||||
// 如果没有人在等待(例如服务重启后),则由本协程负责处理结果入库(记录日志并清理)
|
||||
logger.Infof("[Agent] 没有找到等待任务 #%d 结果的 goroutine,直接处理结果", result.TaskID)
|
||||
logger.Infof("[Agent] 没有找到等待任务 #%s 结果的 goroutine,直接处理结果", result.TaskID)
|
||||
sendStatsService := NewSendStatsService()
|
||||
taskLogService := tasks.NewTaskLogService(sendStatsService)
|
||||
|
||||
@@ -379,7 +383,7 @@ func (s *AgentService) ReportResult(result *models.AgentTaskResult) error {
|
||||
}
|
||||
|
||||
// UpdateTaskDuration 更新任务耗时(心跳)
|
||||
func (s *AgentService) UpdateTaskDuration(logID uint, duration int64) error {
|
||||
func (s *AgentService) UpdateTaskDuration(logID string, duration int64) error {
|
||||
taskLogService := tasks.NewTaskLogService(nil)
|
||||
return taskLogService.UpdateTaskDuration(logID, duration)
|
||||
}
|
||||
@@ -504,12 +508,12 @@ func (s *AgentService) GetAgentBinary(osType, arch string) ([]byte, string, erro
|
||||
}
|
||||
|
||||
// SetForceUpdate 设置强制更新标志
|
||||
func (s *AgentService) SetForceUpdate(id uint) error {
|
||||
func (s *AgentService) SetForceUpdate(id string) error {
|
||||
return database.DB.Model(&models.Agent{}).Where("id = ?", id).Update("force_update", true).Error
|
||||
}
|
||||
|
||||
// ClearForceUpdate 清除强制更新标志
|
||||
func (s *AgentService) ClearForceUpdate(id uint) error {
|
||||
func (s *AgentService) ClearForceUpdate(id string) error {
|
||||
return database.DB.Model(&models.Agent{}).Where("id = ?", id).Update("force_update", false).Error
|
||||
}
|
||||
|
||||
|
||||
@@ -15,11 +15,11 @@ import (
|
||||
|
||||
// AgentWSManager WebSocket 连接管理器
|
||||
type AgentWSManager struct {
|
||||
connections map[uint]*AgentConnection // Agent ID -> 连接对象
|
||||
connections map[string]*AgentConnection // Agent ID -> 连接对象
|
||||
ipConnections map[string]int // IP -> 连接数
|
||||
ipLastAttempt map[string]time.Time // IP -> 最后连接尝试时间
|
||||
ipFailCount map[string]int // IP -> 连续失败次数
|
||||
remoteWaiters map[uint]chan *models.AgentTaskResult // 日志 ID -> 结果通道
|
||||
remoteWaiters map[string]chan *models.AgentTaskResult // 日志 ID -> 结果通道
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
@@ -33,7 +33,7 @@ const (
|
||||
|
||||
// AgentConnection Agent WebSocket 连接
|
||||
type AgentConnection struct {
|
||||
AgentID uint
|
||||
AgentID string
|
||||
IP string
|
||||
Conn *websocket.Conn
|
||||
Send chan []byte
|
||||
@@ -72,11 +72,11 @@ var agentWSOnce sync.Once
|
||||
func GetAgentWSManager() *AgentWSManager {
|
||||
agentWSOnce.Do(func() {
|
||||
agentWSManager = &AgentWSManager{
|
||||
connections: make(map[uint]*AgentConnection),
|
||||
connections: make(map[string]*AgentConnection),
|
||||
ipConnections: make(map[string]int),
|
||||
ipLastAttempt: make(map[string]time.Time),
|
||||
ipFailCount: make(map[string]int),
|
||||
remoteWaiters: make(map[uint]chan *models.AgentTaskResult),
|
||||
remoteWaiters: make(map[string]chan *models.AgentTaskResult),
|
||||
}
|
||||
go agentWSManager.cleanupLoop()
|
||||
})
|
||||
@@ -137,7 +137,7 @@ func (m *AgentWSManager) RecordConnectSuccess(ip string) {
|
||||
}
|
||||
|
||||
// Register 注册连接
|
||||
func (m *AgentWSManager) Register(agentID uint, conn *websocket.Conn, ip string) *AgentConnection {
|
||||
func (m *AgentWSManager) Register(agentID string, conn *websocket.Conn, ip string) *AgentConnection {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
@@ -164,12 +164,12 @@ func (m *AgentWSManager) Register(agentID uint, conn *websocket.Conn, ip string)
|
||||
// 增加 IP 连接计数
|
||||
m.ipConnections[ip]++
|
||||
|
||||
logger.Infof("[AgentWS] Agent #%d 已连接 (%s)", agentID, ip)
|
||||
logger.Infof("[AgentWS] Agent #%s 已连接 (%s)", agentID, ip)
|
||||
return ac
|
||||
}
|
||||
|
||||
// Unregister 注销连接(只注销指定的连接实例)
|
||||
func (m *AgentWSManager) Unregister(agentID uint, ac *AgentConnection) {
|
||||
func (m *AgentWSManager) Unregister(agentID string, ac *AgentConnection) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
@@ -183,19 +183,19 @@ func (m *AgentWSManager) Unregister(agentID uint, ac *AgentConnection) {
|
||||
}
|
||||
conn.Close()
|
||||
delete(m.connections, agentID)
|
||||
logger.Infof("[AgentWS] Agent #%d 已断开", agentID)
|
||||
logger.Infof("[AgentWS] Agent #%s 已断开", agentID)
|
||||
}
|
||||
}
|
||||
|
||||
// GetConnection 获取连接
|
||||
func (m *AgentWSManager) GetConnection(agentID uint) *AgentConnection {
|
||||
func (m *AgentWSManager) GetConnection(agentID string) *AgentConnection {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.connections[agentID]
|
||||
}
|
||||
|
||||
// SendToAgent 发送消息给指定 Agent
|
||||
func (m *AgentWSManager) SendToAgent(agentID uint, msgType string, data interface{}) error {
|
||||
func (m *AgentWSManager) SendToAgent(agentID string, msgType string, data interface{}) error {
|
||||
conn := m.GetConnection(agentID)
|
||||
if conn == nil {
|
||||
return nil // Agent 不在线
|
||||
@@ -214,7 +214,7 @@ func (m *AgentWSManager) SendToAgent(agentID uint, msgType string, data interfac
|
||||
}
|
||||
|
||||
// BroadcastTasks 广播任务更新给指定 Agent
|
||||
func (m *AgentWSManager) BroadcastTasks(agentID uint) {
|
||||
func (m *AgentWSManager) BroadcastTasks(agentID string) {
|
||||
agentService := NewAgentService()
|
||||
tasks := agentService.GetTasks(agentID)
|
||||
m.SendToAgent(agentID, WSTypeTasks, map[string]interface{}{
|
||||
@@ -223,7 +223,7 @@ func (m *AgentWSManager) BroadcastTasks(agentID uint) {
|
||||
}
|
||||
|
||||
// RegisterRemoteWaiter 注册远程任务结果等待者
|
||||
func (m *AgentWSManager) RegisterRemoteWaiter(logID uint) chan *models.AgentTaskResult {
|
||||
func (m *AgentWSManager) RegisterRemoteWaiter(logID string) chan *models.AgentTaskResult {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
ch := make(chan *models.AgentTaskResult, 1)
|
||||
@@ -232,7 +232,7 @@ func (m *AgentWSManager) RegisterRemoteWaiter(logID uint) chan *models.AgentTask
|
||||
}
|
||||
|
||||
// UnregisterRemoteWaiter 注销远程任务结果等待者
|
||||
func (m *AgentWSManager) UnregisterRemoteWaiter(logID uint) {
|
||||
func (m *AgentWSManager) UnregisterRemoteWaiter(logID string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.remoteWaiters, logID)
|
||||
@@ -293,7 +293,7 @@ func (m *AgentWSManager) cleanupLoop() {
|
||||
delete(m.connections, agentID)
|
||||
// 更新数据库状态
|
||||
database.DB.Model(&models.Agent{}).Where("id = ?", agentID).Update("status", constant.AgentStatusOffline)
|
||||
logger.Infof("[AgentWS] Agent #%d 心跳超时,已断开", agentID)
|
||||
logger.Infof("[AgentWS] Agent #%s 心跳超时,已断开", agentID)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -158,7 +158,7 @@ func (s *BackupService) CreateBackup() (string, error) {
|
||||
|
||||
// 写入元数据信息
|
||||
sysInfo := map[string]interface{}{
|
||||
"version": "v2",
|
||||
"version": "v3",
|
||||
"ts": time.Now().Format("2006-01-02 15:04:05"),
|
||||
}
|
||||
sysFile, err := zipWriter.Create("__sys__.json")
|
||||
@@ -197,6 +197,24 @@ func (s *BackupService) Restore(zipPath string) error {
|
||||
fileMap[f.Name] = f
|
||||
}
|
||||
|
||||
// 校验版本
|
||||
if f, ok := fileMap["__sys__.json"]; ok {
|
||||
rc, err := f.Open()
|
||||
if err == nil {
|
||||
var sysInfo map[string]interface{}
|
||||
json.NewDecoder(rc).Decode(&sysInfo)
|
||||
rc.Close()
|
||||
if v, ok := sysInfo["version"]; ok {
|
||||
vs, _ := v.(string)
|
||||
if vs < "v3" {
|
||||
return fmt.Errorf("只能数据随版本升级上来,当前备份版本为 %s,限制 v3 以下的不能导入", vs)
|
||||
}
|
||||
}
|
||||
}
|
||||
// } else {
|
||||
// return fmt.Errorf("非法备份包:缺失版本标记")
|
||||
}
|
||||
|
||||
// 开启全局事务
|
||||
return database.DB.Transaction(func(tx *gorm.DB) error {
|
||||
// 1. 清空现有数据(物理删除)
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"github.com/engigu/baihu-panel/internal/database"
|
||||
"github.com/engigu/baihu-panel/internal/models"
|
||||
"github.com/engigu/baihu-panel/internal/services/deps"
|
||||
"github.com/engigu/baihu-panel/internal/utils"
|
||||
)
|
||||
|
||||
type DependencyService struct{}
|
||||
@@ -36,12 +37,15 @@ func (s *DependencyService) Create(dep *models.Dependency) error {
|
||||
if err == nil {
|
||||
return errors.New("依赖已存在")
|
||||
}
|
||||
if dep.ID == "" {
|
||||
dep.ID = utils.GenerateID()
|
||||
}
|
||||
return database.DB.Create(dep).Error
|
||||
}
|
||||
|
||||
// Delete 删除依赖记录
|
||||
func (s *DependencyService) Delete(id int) error {
|
||||
return database.DB.Delete(&models.Dependency{}, id).Error
|
||||
func (s *DependencyService) Delete(id string) error {
|
||||
return database.DB.Where("id = ?", id).Delete(&models.Dependency{}).Error
|
||||
}
|
||||
|
||||
// Install 安装依赖
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/baihu-panel/internal/database"
|
||||
"github.com/engigu/baihu-panel/internal/models"
|
||||
"github.com/engigu/baihu-panel/internal/utils"
|
||||
)
|
||||
|
||||
type EnvService struct{}
|
||||
@@ -14,42 +14,26 @@ func NewEnvService() *EnvService {
|
||||
return &EnvService{}
|
||||
}
|
||||
|
||||
func (es *EnvService) CreateEnvVar(name, value, remark string, hidden bool, userID int) *models.EnvironmentVariable {
|
||||
func (es *EnvService) CreateEnvVar(name, value, remark string, hidden bool, userID string) *models.EnvironmentVariable {
|
||||
env := &models.EnvironmentVariable{
|
||||
ID: utils.GenerateID(),
|
||||
Name: name,
|
||||
Value: value,
|
||||
Remark: remark,
|
||||
Hidden: hidden,
|
||||
UserID: uint(userID),
|
||||
}
|
||||
data := map[string]interface{}{
|
||||
"name": name,
|
||||
"value": value,
|
||||
"remark": remark,
|
||||
"hidden": hidden,
|
||||
"user_id": userID,
|
||||
}
|
||||
database.DB.Model(&models.EnvironmentVariable{}).Create(data)
|
||||
|
||||
// 将自动生成的 ID 赋值回对象以便返回
|
||||
if id, ok := data["id"].(uint); ok {
|
||||
env.ID = id
|
||||
} else if id, ok := data["id"].(int64); ok {
|
||||
env.ID = uint(id)
|
||||
} else {
|
||||
// 如果 ID 没有自动回填到 map,尝试通过刚才的数据查出来
|
||||
database.DB.Where("name = ? AND user_id = ?", name, userID).Order("id DESC").First(env)
|
||||
UserID: userID,
|
||||
}
|
||||
database.DB.Create(env)
|
||||
return env
|
||||
}
|
||||
|
||||
func (es *EnvService) GetEnvVarsByUserID(userID int) []models.EnvironmentVariable {
|
||||
func (es *EnvService) GetEnvVarsByUserID(userID string) []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) {
|
||||
func (es *EnvService) GetEnvVarsWithPagination(userID string, name string, page, pageSize int) ([]models.EnvironmentVariable, int64) {
|
||||
var envs []models.EnvironmentVariable
|
||||
var total int64
|
||||
|
||||
@@ -63,17 +47,17 @@ func (es *EnvService) GetEnvVarsWithPagination(userID int, name string, page, pa
|
||||
return envs, total
|
||||
}
|
||||
|
||||
func (es *EnvService) GetEnvVarByID(id int) *models.EnvironmentVariable {
|
||||
func (es *EnvService) GetEnvVarByID(id string) *models.EnvironmentVariable {
|
||||
var env models.EnvironmentVariable
|
||||
if err := database.DB.First(&env, id).Error; err != nil {
|
||||
if err := database.DB.Where("id = ?", id).First(&env).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return &env
|
||||
}
|
||||
|
||||
func (es *EnvService) UpdateEnvVar(id int, name, value, remark string, hidden bool) *models.EnvironmentVariable {
|
||||
func (es *EnvService) UpdateEnvVar(id string, name, value, remark string, hidden bool) *models.EnvironmentVariable {
|
||||
var env models.EnvironmentVariable
|
||||
if err := database.DB.First(&env, id).Error; err != nil {
|
||||
if err := database.DB.Where("id = ?", id).First(&env).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
env.Name = name
|
||||
@@ -84,8 +68,8 @@ func (es *EnvService) UpdateEnvVar(id int, name, value, remark string, hidden bo
|
||||
return &env
|
||||
}
|
||||
|
||||
func (es *EnvService) DeleteEnvVar(id int) bool {
|
||||
result := database.DB.Delete(&models.EnvironmentVariable{}, id)
|
||||
func (es *EnvService) DeleteEnvVar(id string) bool {
|
||||
result := database.DB.Where("id = ?", id).Delete(&models.EnvironmentVariable{})
|
||||
return result.RowsAffected > 0
|
||||
}
|
||||
|
||||
@@ -107,12 +91,12 @@ func (es *EnvService) GetEnvVarsByIDs(envIDs string) []string {
|
||||
}
|
||||
|
||||
// splitEnvIDs 解析逗号分隔的ID字符串
|
||||
func splitEnvIDs(envIDs string) []int {
|
||||
var ids []int
|
||||
func splitEnvIDs(envIDs string) []string {
|
||||
var ids []string
|
||||
for _, s := range strings.Split(envIDs, ",") {
|
||||
s = strings.TrimSpace(s)
|
||||
if id, err := strconv.Atoi(s); err == nil {
|
||||
ids = append(ids, id)
|
||||
if s != "" {
|
||||
ids = append(ids, s)
|
||||
}
|
||||
}
|
||||
return ids
|
||||
|
||||
@@ -3,6 +3,7 @@ package services
|
||||
import (
|
||||
"github.com/engigu/baihu-panel/internal/database"
|
||||
"github.com/engigu/baihu-panel/internal/models"
|
||||
"github.com/engigu/baihu-panel/internal/utils"
|
||||
)
|
||||
|
||||
type LoginLogService struct{}
|
||||
@@ -14,6 +15,7 @@ func NewLoginLogService() *LoginLogService {
|
||||
// Create 创建登录日志
|
||||
func (s *LoginLogService) Create(username, ip, userAgent, status, message string) error {
|
||||
log := &models.LoginLog{
|
||||
ID: utils.GenerateID(),
|
||||
Username: username,
|
||||
IP: ip,
|
||||
UserAgent: userAgent,
|
||||
|
||||
@@ -0,0 +1,326 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/baihu-panel/internal/database"
|
||||
"github.com/engigu/baihu-panel/internal/logger"
|
||||
"github.com/engigu/baihu-panel/internal/models"
|
||||
"github.com/engigu/baihu-panel/internal/utils"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// MigrationTable 定义迁移配置
|
||||
type MigrationTable struct {
|
||||
Model any
|
||||
EntityName string
|
||||
FKs map[string]string // 单个外键映射: 字段名 -> 实体名
|
||||
MultiFKs map[string]string // 复合外键映射 (如逗号分隔): 字段名 -> 实体名
|
||||
}
|
||||
|
||||
func getMigrationTables() []MigrationTable {
|
||||
return []MigrationTable{
|
||||
{&models.User{}, "users", nil, nil},
|
||||
{&models.Agent{}, "agents", nil, nil},
|
||||
{&models.AgentToken{}, "tokens", nil, nil},
|
||||
{&models.EnvironmentVariable{}, "envs", map[string]string{"UserID": "users"}, nil},
|
||||
{&models.Task{}, "tasks", map[string]string{"AgentID": "agents"}, map[string]string{"Envs": "envs"}},
|
||||
{&models.TaskLog{}, "task_logs", map[string]string{"TaskID": "tasks", "AgentID": "agents"}, nil},
|
||||
{&models.Script{}, "scripts", map[string]string{"UserID": "users"}, nil},
|
||||
{&models.Setting{}, "settings", nil, nil},
|
||||
{&models.SendStats{}, "send_stats", map[string]string{"TaskID": "tasks"}, nil},
|
||||
{&models.LoginLog{}, "login_logs", nil, nil},
|
||||
{&models.Language{}, "languages", nil, nil},
|
||||
{&models.Dependency{}, "deps", nil, nil},
|
||||
}
|
||||
}
|
||||
|
||||
func getTableName(db *gorm.DB, model any) string {
|
||||
stmt := &gorm.Statement{DB: db}
|
||||
if err := stmt.Parse(model); err != nil {
|
||||
return ""
|
||||
}
|
||||
return stmt.Schema.Table
|
||||
}
|
||||
|
||||
func isTableStringID(db *gorm.DB, model any) bool {
|
||||
if !db.Migrator().HasTable(model) {
|
||||
return true
|
||||
}
|
||||
columnTypes, err := db.Migrator().ColumnTypes(model)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
for _, ct := range columnTypes {
|
||||
if strings.ToLower(ct.Name()) == "id" {
|
||||
typeName := strings.ToLower(ct.DatabaseTypeName())
|
||||
return strings.Contains(typeName, "char") || strings.Contains(typeName, "text") || strings.Contains(typeName, "string")
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func getValFromMap(m map[string]interface{}, key string) (interface{}, bool) {
|
||||
lowerKey := strings.ToLower(key)
|
||||
for k, v := range m {
|
||||
if strings.ToLower(k) == lowerKey {
|
||||
return v, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func RunMigrationV3() error {
|
||||
db := database.DB
|
||||
if db == nil {
|
||||
return fmt.Errorf("数据库未初始化")
|
||||
}
|
||||
|
||||
// 0. 检查迁移标记,防止重复迁移逻辑被误判触发
|
||||
if db.Migrator().HasTable(&models.Setting{}) {
|
||||
var migrationFlag models.Setting
|
||||
err := db.Where("section = ? AND `key` = ?", "system", "migration_v3_success").First(&migrationFlag).Error
|
||||
if err == nil && migrationFlag.Value == "true" {
|
||||
// 如果已经是字符串 ID 模式,双重确认
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
tables := getMigrationTables()
|
||||
needMigration := false
|
||||
for _, t := range tables {
|
||||
if !isTableStringID(db, t.Model) {
|
||||
needMigration = true
|
||||
logger.Infof("[MigrationV3] 表 [%s] ID 为数字,需迁移", getTableName(db, t.Model))
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !needMigration {
|
||||
// 没有数字 ID 表了,但也可能还没打标(比如之前手动修过表),此时补一个标
|
||||
return markMigrationSuccess(db)
|
||||
}
|
||||
|
||||
// 1. 备份过程
|
||||
backupDir := "./data"
|
||||
os.MkdirAll(backupDir, 0755)
|
||||
// 如果已经有还原过的记录,为了防止反复循环,我们可以检查还原标记
|
||||
backups, _ := filepath.Glob(filepath.Join(backupDir, "migration_v3_backup_*.zip"))
|
||||
if len(backups) == 0 {
|
||||
logger.Infof("[MigrationV3] 执行关键备份...")
|
||||
backupService := NewBackupService()
|
||||
zipPath, err := backupService.CreateBackup()
|
||||
if err != nil {
|
||||
return fmt.Errorf("自动备份失败,流程终止: %v", err)
|
||||
}
|
||||
newPath := filepath.Join(backupDir, fmt.Sprintf("migration_v3_backup_%s.zip", filepath.Base(zipPath)))
|
||||
os.Rename(zipPath, newPath)
|
||||
logger.Infof("[MigrationV3] 备份成功: %s", newPath)
|
||||
}
|
||||
|
||||
mappings := make(map[string]map[uint]string)
|
||||
err := db.Transaction(func(tx *gorm.DB) error {
|
||||
return performHardMigration(tx, mappings)
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 3. 标记成功
|
||||
return markMigrationSuccess(db)
|
||||
}
|
||||
|
||||
func markMigrationSuccess(db *gorm.DB) error {
|
||||
if !db.Migrator().HasTable(&models.Setting{}) {
|
||||
return nil
|
||||
}
|
||||
var flag models.Setting
|
||||
err := db.Where("section = ? AND `key` = ?", "system", "migration_v3_success").First(&flag).Error
|
||||
if err != nil {
|
||||
// 创建或更新
|
||||
flag = models.Setting{
|
||||
ID: utils.GenerateID(),
|
||||
Section: "system",
|
||||
Key: "migration_v3_success",
|
||||
Value: "true",
|
||||
}
|
||||
return db.Create(&flag).Error
|
||||
}
|
||||
return db.Model(&flag).Update("value", "true").Error
|
||||
}
|
||||
|
||||
func performHardMigration(tx *gorm.DB, mappings map[string]map[uint]string) error {
|
||||
allTables := getMigrationTables()
|
||||
|
||||
// ---------------------------------------------------------
|
||||
// 第一阶段:全量构建 ID 映射映射表 (Pass 1)
|
||||
// ---------------------------------------------------------
|
||||
for _, t := range allTables {
|
||||
actualName := getTableName(tx, t.Model)
|
||||
if actualName == "" || !tx.Migrator().HasTable(actualName) {
|
||||
continue
|
||||
}
|
||||
mappings[t.EntityName] = make(map[uint]string)
|
||||
oldTableName := actualName + "_v2_bak"
|
||||
|
||||
// 如果还没有备份表,说明这是第一次处理该表,先重命名
|
||||
if !tx.Migrator().HasTable(oldTableName) {
|
||||
if isTableStringID(tx, t.Model) {
|
||||
continue // 已经是字符串 ID 且无备份,跳过
|
||||
}
|
||||
if err := tx.Migrator().RenameTable(actualName, oldTableName); err != nil {
|
||||
return fmt.Errorf("重命名表 %s 失败: %v", actualName, err)
|
||||
}
|
||||
}
|
||||
|
||||
// 预先为该表所有记录生成新的 xid
|
||||
var rows []map[string]interface{}
|
||||
tx.Table(oldTableName).Select("id").Find(&rows)
|
||||
for _, row := range rows {
|
||||
if val, ok := getValFromMap(row, "id"); ok {
|
||||
uid := parseUint(val)
|
||||
if uid > 0 {
|
||||
mappings[t.EntityName][uid] = utils.GenerateID()
|
||||
}
|
||||
}
|
||||
}
|
||||
logger.Infof("[MigrationV3] Pass 1: 构建关键表 %s 的 ID 映射, 共 %d 条", actualName, len(mappings[t.EntityName]))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------
|
||||
// 第二、三阶段:正式转换数据并处理关联字段 (Pass 2 & 3)
|
||||
// ---------------------------------------------------------
|
||||
for _, t := range allTables {
|
||||
actualName := getTableName(tx, t.Model)
|
||||
oldTableName := actualName + "_v2_bak"
|
||||
|
||||
if !tx.Migrator().HasTable(oldTableName) {
|
||||
// 虽然可能已经改过格式,但为了安全还是 AutoMigrate 一下
|
||||
tx.AutoMigrate(t.Model)
|
||||
continue
|
||||
}
|
||||
|
||||
logger.Infof("[MigrationV3] Pass 2&3: 正在转换数据并修复关联: %s", actualName)
|
||||
tx.AutoMigrate(t.Model)
|
||||
|
||||
// 获取新表的有效列名(小写)
|
||||
columnTypes, _ := tx.Migrator().ColumnTypes(t.Model)
|
||||
validColumns := make(map[string]bool)
|
||||
for _, ct := range columnTypes {
|
||||
validColumns[strings.ToLower(ct.Name())] = true
|
||||
}
|
||||
|
||||
var oldData []map[string]interface{}
|
||||
if err := tx.Table(oldTableName).Find(&oldData).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, row := range oldData {
|
||||
// 1. 处理主键 ID
|
||||
if val, ok := getValFromMap(row, "id"); ok {
|
||||
uid := parseUint(val)
|
||||
if nid, exists := mappings[t.EntityName][uid]; exists {
|
||||
row["id"] = nid
|
||||
} else {
|
||||
row["id"] = utils.GenerateID()
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 处理单外键关联 (Phase 3)
|
||||
for field, parentEntity := range t.FKs {
|
||||
columnName := getColumnName(field)
|
||||
if val, ok := getValFromMap(row, columnName); ok && val != nil {
|
||||
ufk := parseUint(val)
|
||||
if ufk > 0 {
|
||||
if nid, exists := mappings[parentEntity][ufk]; exists {
|
||||
row[columnName] = nid
|
||||
} else {
|
||||
row[columnName] = nil
|
||||
}
|
||||
} else {
|
||||
row[columnName] = nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 处理复合多外键字段 (envs)
|
||||
for field, parentEntity := range t.MultiFKs {
|
||||
columnName := getColumnName(field)
|
||||
if val, ok := getValFromMap(row, columnName); ok && val != nil {
|
||||
if strVal, ok := val.(string); ok && strVal != "" {
|
||||
row[columnName] = transformMultiIDs(strVal, parentEntity, mappings)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 4. 过滤不存在的列并插入
|
||||
filteredRow := make(map[string]interface{})
|
||||
for k, v := range row {
|
||||
if validColumns[strings.ToLower(k)] {
|
||||
filteredRow[k] = v
|
||||
}
|
||||
}
|
||||
if err := tx.Table(actualName).Create(filteredRow).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 迁移完成,清理备份表
|
||||
tx.Migrator().DropTable(oldTableName)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// 辅助函数:解析各种数字 ID
|
||||
func parseUint(val interface{}) uint {
|
||||
if val == nil { return 0 }
|
||||
switch v := val.(type) {
|
||||
case uint: return v
|
||||
case int64: return uint(v)
|
||||
case int: return uint(v)
|
||||
case uint64: return uint(v)
|
||||
case float64: return uint(v)
|
||||
case string:
|
||||
var u uint
|
||||
fmt.Sscanf(v, "%d", &u)
|
||||
return u
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// 辅助函数:字段名转列名
|
||||
func getColumnName(field string) string {
|
||||
switch field {
|
||||
case "AgentID": return "agent_id"
|
||||
case "TaskID": return "task_id"
|
||||
case "UserID": return "user_id"
|
||||
case "LogID": return "log_id"
|
||||
case "Envs": return "envs"
|
||||
default: return strings.ToLower(field)
|
||||
}
|
||||
}
|
||||
|
||||
// 辅助函数:处理逗号分隔的 ID 列表
|
||||
func transformMultiIDs(oldStr string, parentEntity string, mappings map[string]map[uint]string) string {
|
||||
parts := strings.Split(oldStr, ",")
|
||||
var result []string
|
||||
for _, p := range parts {
|
||||
p = strings.TrimSpace(p)
|
||||
if p == "" { continue }
|
||||
if len(p) == 20 && !utils.IsNumeric(p) {
|
||||
result = append(result, p) // 已经是 xid,保留
|
||||
continue
|
||||
}
|
||||
uid := parseUint(p)
|
||||
if uid > 0 {
|
||||
if nid, exists := mappings[parentEntity][uid]; exists {
|
||||
result = append(result, nid)
|
||||
}
|
||||
}
|
||||
}
|
||||
return strings.Join(result, ",")
|
||||
}
|
||||
@@ -283,7 +283,7 @@ func (s *MiseService) syncToDB(languages []MiseLanguage) {
|
||||
return
|
||||
}
|
||||
|
||||
var currentIds []uint
|
||||
var currentIds []string
|
||||
for _, lang := range languages {
|
||||
var model models.Language
|
||||
// 以 plugin 和 version 作为联合唯一标识(业务逻辑上)
|
||||
@@ -308,6 +308,7 @@ func (s *MiseService) syncToDB(languages []MiseLanguage) {
|
||||
if err != nil {
|
||||
// 如果不存在,则创建
|
||||
newLang := models.Language{
|
||||
ID: utils.GenerateID(),
|
||||
Plugin: lang.Plugin,
|
||||
Version: lang.Version,
|
||||
InstallPath: lang.InstallPath,
|
||||
|
||||
@@ -3,6 +3,7 @@ package services
|
||||
import (
|
||||
"github.com/engigu/baihu-panel/internal/database"
|
||||
"github.com/engigu/baihu-panel/internal/models"
|
||||
"github.com/engigu/baihu-panel/internal/utils"
|
||||
)
|
||||
|
||||
type ScriptService struct{}
|
||||
@@ -11,33 +12,34 @@ func NewScriptService() *ScriptService {
|
||||
return &ScriptService{}
|
||||
}
|
||||
|
||||
func (ss *ScriptService) CreateScript(name, content string, userID int) *models.Script {
|
||||
func (ss *ScriptService) CreateScript(name, content string, userID string) *models.Script {
|
||||
script := &models.Script{
|
||||
ID: utils.GenerateID(),
|
||||
Name: name,
|
||||
Content: content,
|
||||
UserID: uint(userID),
|
||||
UserID: userID,
|
||||
}
|
||||
database.DB.Create(script)
|
||||
return script
|
||||
}
|
||||
|
||||
func (ss *ScriptService) GetScriptsByUserID(userID int) []models.Script {
|
||||
func (ss *ScriptService) GetScriptsByUserID(userID string) []models.Script {
|
||||
var scripts []models.Script
|
||||
database.DB.Where("user_id = ?", userID).Find(&scripts)
|
||||
return scripts
|
||||
}
|
||||
|
||||
func (ss *ScriptService) GetScriptByID(id int) *models.Script {
|
||||
func (ss *ScriptService) GetScriptByID(id string) *models.Script {
|
||||
var script models.Script
|
||||
if err := database.DB.First(&script, id).Error; err != nil {
|
||||
if err := database.DB.Where("id = ?", id).First(&script).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return &script
|
||||
}
|
||||
|
||||
func (ss *ScriptService) UpdateScript(id int, name, content string) *models.Script {
|
||||
func (ss *ScriptService) UpdateScript(id string, name, content string) *models.Script {
|
||||
var script models.Script
|
||||
if err := database.DB.First(&script, id).Error; err != nil {
|
||||
if err := database.DB.Where("id = ?", id).First(&script).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
script.Name = name
|
||||
@@ -46,7 +48,7 @@ func (ss *ScriptService) UpdateScript(id int, name, content string) *models.Scri
|
||||
return &script
|
||||
}
|
||||
|
||||
func (ss *ScriptService) DeleteScript(id int) bool {
|
||||
result := database.DB.Delete(&models.Script{}, id)
|
||||
func (ss *ScriptService) DeleteScript(id string) bool {
|
||||
result := database.DB.Where("id = ?", id).Delete(&models.Script{})
|
||||
return result.RowsAffected > 0
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"github.com/engigu/baihu-panel/internal/database"
|
||||
"github.com/engigu/baihu-panel/internal/models"
|
||||
"github.com/engigu/baihu-panel/internal/systime"
|
||||
"github.com/engigu/baihu-panel/internal/utils"
|
||||
)
|
||||
|
||||
type SendStatsService struct{}
|
||||
@@ -15,7 +16,7 @@ func NewSendStatsService() *SendStatsService {
|
||||
}
|
||||
|
||||
// IncrementStats 增加任务执行统计
|
||||
func (s *SendStatsService) IncrementStats(taskID uint, status string) error {
|
||||
func (s *SendStatsService) IncrementStats(taskID string, status string) error {
|
||||
day := systime.FormatDate(time.Now())
|
||||
|
||||
var stats models.SendStats
|
||||
@@ -24,6 +25,7 @@ func (s *SendStatsService) IncrementStats(taskID uint, status string) error {
|
||||
if result.Error != nil {
|
||||
// 不存在则创建
|
||||
stats = models.SendStats{
|
||||
ID: utils.GenerateID(),
|
||||
TaskID: taskID,
|
||||
Day: day,
|
||||
Status: status,
|
||||
@@ -37,7 +39,7 @@ func (s *SendStatsService) IncrementStats(taskID uint, status string) error {
|
||||
}
|
||||
|
||||
// GetStatsByTaskID 获取任务的统计数据
|
||||
func (s *SendStatsService) GetStatsByTaskID(taskID uint) []models.SendStats {
|
||||
func (s *SendStatsService) GetStatsByTaskID(taskID string) []models.SendStats {
|
||||
var stats []models.SendStats
|
||||
database.DB.Where("task_id = ?", taskID).Order("day DESC").Find(&stats)
|
||||
return stats
|
||||
|
||||
@@ -21,7 +21,12 @@ func (s *SettingsService) InitSettings() error {
|
||||
var count int64
|
||||
database.DB.Model(&models.Setting{}).Where("section = ? AND `key` = ?", section, key).Count(&count)
|
||||
if count == 0 {
|
||||
if err := database.DB.Create(&models.Setting{Section: section, Key: key, Value: value}).Error; err != nil {
|
||||
if err := database.DB.Create(&models.Setting{
|
||||
ID: utils.GenerateID(),
|
||||
Section: section,
|
||||
Key: key,
|
||||
Value: value,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
@@ -38,7 +43,12 @@ func (s *SettingsService) InitSettings() error {
|
||||
} else {
|
||||
secretValue = utils.RandomString(32)
|
||||
}
|
||||
if err := database.DB.Create(&models.Setting{Section: constant.SectionSecurity, Key: constant.KeySecret, Value: secretValue}).Error; err != nil {
|
||||
if err := database.DB.Create(&models.Setting{
|
||||
ID: utils.GenerateID(),
|
||||
Section: constant.SectionSecurity,
|
||||
Key: constant.KeySecret,
|
||||
Value: secretValue,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
} else {
|
||||
@@ -69,7 +79,12 @@ func (s *SettingsService) Get(section, key string) string {
|
||||
func (s *SettingsService) Set(section, key, value string) error {
|
||||
var setting models.Setting
|
||||
if database.DB.Where("section = ? AND `key` = ?", section, key).First(&setting).Error != nil {
|
||||
return database.DB.Create(&models.Setting{Section: section, Key: key, Value: value}).Error
|
||||
return database.DB.Create(&models.Setting{
|
||||
ID: utils.GenerateID(),
|
||||
Section: section,
|
||||
Key: key,
|
||||
Value: value,
|
||||
}).Error
|
||||
}
|
||||
return database.DB.Model(&setting).Update("value", value).Error
|
||||
}
|
||||
|
||||
@@ -24,9 +24,9 @@ import (
|
||||
|
||||
// AgentWSManager 接口定义(避免循环依赖)
|
||||
type AgentWSManager interface {
|
||||
RegisterRemoteWaiter(logID uint) chan *models.AgentTaskResult
|
||||
UnregisterRemoteWaiter(logID uint)
|
||||
SendToAgent(agentID uint, msgType string, data interface{}) error
|
||||
RegisterRemoteWaiter(logID string) chan *models.AgentTaskResult
|
||||
UnregisterRemoteWaiter(logID string)
|
||||
SendToAgent(agentID string, msgType string, data interface{}) error
|
||||
}
|
||||
|
||||
// SettingsService 接口定义(避免循环依赖)
|
||||
@@ -115,10 +115,9 @@ func (h *ServerSchedulerHandler) OnTaskScheduled(req *executor.ExecutionRequest)
|
||||
}
|
||||
|
||||
func (h *ServerSchedulerHandler) OnTaskExecuting(req *executor.ExecutionRequest) (io.Writer, io.Writer, error) {
|
||||
var taskID uint
|
||||
fmt.Sscanf(req.TaskID, "%d", &taskID)
|
||||
taskID := req.TaskID
|
||||
|
||||
task := h.es.taskService.GetTaskByID(int(taskID))
|
||||
task := h.es.taskService.GetTaskByID(taskID)
|
||||
// 系统任务(无 taskID)不记录数据库日志,直接返回空写入器
|
||||
if task == nil {
|
||||
return nil, nil, nil
|
||||
@@ -168,7 +167,7 @@ func (h *ServerSchedulerHandler) OnTaskExecuting(req *executor.ExecutionRequest)
|
||||
}
|
||||
|
||||
func (h *ServerSchedulerHandler) OnTaskHeartbeat(req *executor.ExecutionRequest, duration int64) {
|
||||
if req.LogID > 0 {
|
||||
if req.LogID != "" {
|
||||
h.es.taskLogService.UpdateTaskDuration(req.LogID, duration)
|
||||
}
|
||||
|
||||
@@ -184,14 +183,13 @@ func (h *ServerSchedulerHandler) OnTaskStarted(req *executor.ExecutionRequest) {
|
||||
}
|
||||
|
||||
func (h *ServerSchedulerHandler) OnTaskCompleted(req *executor.ExecutionRequest, result *executor.ExecutionResult) {
|
||||
if req.LogID == 0 {
|
||||
if req.LogID == "" {
|
||||
return
|
||||
}
|
||||
|
||||
var taskID uint
|
||||
fmt.Sscanf(req.TaskID, "%d", &taskID)
|
||||
taskID := req.TaskID
|
||||
|
||||
task := h.es.taskService.GetTaskByID(int(taskID))
|
||||
task := h.es.taskService.GetTaskByID(taskID)
|
||||
if task == nil {
|
||||
return
|
||||
}
|
||||
@@ -204,7 +202,7 @@ func (h *ServerSchedulerHandler) OnTaskCompleted(req *executor.ExecutionRequest,
|
||||
var err error
|
||||
output, err = tl.CompressAndCleanup()
|
||||
if err != nil {
|
||||
logger.Errorf("[Executor] 压缩任务 #%d 日志失败: %v", task.ID, err)
|
||||
logger.Errorf("[Executor] 压缩任务 #%s 日志失败: %v", task.ID, err)
|
||||
output = "[System Error] 日志处理失败: " + err.Error()
|
||||
}
|
||||
} else {
|
||||
@@ -230,7 +228,7 @@ func (h *ServerSchedulerHandler) OnTaskCompleted(req *executor.ExecutionRequest,
|
||||
}
|
||||
|
||||
// 如果有 AgentID,也记录下来
|
||||
if task.AgentID != nil && *task.AgentID > 0 {
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
agentID := *task.AgentID
|
||||
taskLog.AgentID = &agentID
|
||||
}
|
||||
@@ -251,12 +249,11 @@ func (h *ServerSchedulerHandler) OnTaskCompleted(req *executor.ExecutionRequest,
|
||||
}
|
||||
|
||||
func (h *ServerSchedulerHandler) OnTaskFailed(req *executor.ExecutionRequest, err error) {
|
||||
if req.LogID == 0 {
|
||||
if req.LogID == "" {
|
||||
return
|
||||
}
|
||||
|
||||
var taskID uint
|
||||
fmt.Sscanf(req.TaskID, "%d", &taskID)
|
||||
taskID := req.TaskID
|
||||
|
||||
// 移除运行记录
|
||||
if req.Metadata.GoID != 0 {
|
||||
@@ -288,8 +285,8 @@ func (h *ServerSchedulerHandler) OnTaskFailed(req *executor.ExecutionRequest, er
|
||||
}
|
||||
|
||||
// 补充 AgentID
|
||||
task := h.es.taskService.GetTaskByID(int(taskID))
|
||||
if task != nil && task.AgentID != nil && *task.AgentID > 0 {
|
||||
task := h.es.taskService.GetTaskByID(taskID)
|
||||
if task != nil && task.AgentID != nil && *task.AgentID != "" {
|
||||
agentID := *task.AgentID
|
||||
taskLog.AgentID = &agentID
|
||||
}
|
||||
@@ -321,10 +318,10 @@ func (es *ExecutorService) HandleTaskRetry(task *models.Task, req *executor.Exec
|
||||
|
||||
if retryIndex < task.RetryCount {
|
||||
retryIndex++
|
||||
logger.Infof("[Executor] 任务 #%d 执行失败/出错,将在 %d 秒后进行第 %d/%d 次重试...", task.ID, task.RetryInterval, retryIndex, task.RetryCount)
|
||||
logger.Infof("[Executor] 任务 #%s 执行失败/出错,将在 %d 秒后进行第 %d/%d 次重试...", task.ID, task.RetryInterval, retryIndex, task.RetryCount)
|
||||
|
||||
es.scheduler.EnqueueDelayed(time.Duration(task.RetryInterval)*time.Second, func() *executor.ExecutionRequest {
|
||||
latestTask := es.taskService.GetTaskByID(int(task.ID))
|
||||
latestTask := es.taskService.GetTaskByID(task.ID)
|
||||
if latestTask == nil || !latestTask.Enabled {
|
||||
return nil
|
||||
}
|
||||
@@ -350,8 +347,7 @@ func (es *ExecutorService) HandleTaskRetry(task *models.Task, req *executor.Exec
|
||||
}
|
||||
|
||||
func (h *ServerSchedulerHandler) OnCronNextRun(req *executor.ExecutionRequest, nextRun time.Time) {
|
||||
var taskID uint
|
||||
fmt.Sscanf(req.TaskID, "%d", &taskID)
|
||||
taskID := req.TaskID
|
||||
// 更新数据库中的下次运行时间
|
||||
database.DB.Model(&models.Task{}).Where("id = ?", taskID).Update("next_run", nextRun)
|
||||
}
|
||||
@@ -359,19 +355,19 @@ func (h *ServerSchedulerHandler) OnCronNextRun(req *executor.ExecutionRequest, n
|
||||
// LocalTaskHooks 本地任务钩子适配器
|
||||
type LocalTaskHooks struct {
|
||||
es *ExecutorService
|
||||
logID uint
|
||||
logID string
|
||||
}
|
||||
|
||||
func (h *LocalTaskHooks) PreExecute(ctx context.Context, req executor.Request) (uint, error) {
|
||||
func (h *LocalTaskHooks) PreExecute(ctx context.Context, req executor.Request) (string, error) {
|
||||
return h.logID, nil
|
||||
}
|
||||
|
||||
func (h *LocalTaskHooks) PostExecute(ctx context.Context, logID uint, result *executor.Result) error {
|
||||
func (h *LocalTaskHooks) PostExecute(ctx context.Context, logID string, result *executor.Result) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *LocalTaskHooks) OnHeartbeat(ctx context.Context, logID uint, duration int64) error {
|
||||
if logID > 0 {
|
||||
func (h *LocalTaskHooks) OnHeartbeat(ctx context.Context, logID string, duration int64) error {
|
||||
if logID != "" {
|
||||
return h.es.taskLogService.UpdateTaskDuration(logID, duration)
|
||||
}
|
||||
return nil
|
||||
@@ -379,10 +375,9 @@ func (h *LocalTaskHooks) OnHeartbeat(ctx context.Context, logID uint, duration i
|
||||
|
||||
// ExecuteDispatcher 实现任务分发逻辑
|
||||
func (es *ExecutorService) ExecuteDispatcher(ctx context.Context, req *executor.ExecutionRequest, stdout, stderr io.Writer) (*executor.Result, error) {
|
||||
var taskID uint
|
||||
fmt.Sscanf(req.TaskID, "%d", &taskID)
|
||||
taskID := req.TaskID
|
||||
|
||||
task := es.taskService.GetTaskByID(int(taskID))
|
||||
task := es.taskService.GetTaskByID(taskID)
|
||||
// 系统任务(无 taskID)直接本地执行
|
||||
if task == nil {
|
||||
return executor.Execute(ctx, executor.Request{
|
||||
@@ -409,7 +404,7 @@ func (es *ExecutorService) ExecuteDispatcher(ctx context.Context, req *executor.
|
||||
}
|
||||
|
||||
// 远程任务
|
||||
if task.AgentID != nil && *task.AgentID > 0 {
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
return es.ExecuteRemoteForScheduler(task, req.LogID)
|
||||
}
|
||||
|
||||
@@ -467,8 +462,8 @@ func (es *ExecutorService) AddCronTask(task *models.Task) error {
|
||||
}
|
||||
|
||||
// RemoveCronTask 移除计划任务
|
||||
func (es *ExecutorService) RemoveCronTask(taskID uint) {
|
||||
es.cronManager.RemoveTask(fmt.Sprintf("%d", taskID))
|
||||
func (es *ExecutorService) RemoveCronTask(taskID string) {
|
||||
es.cronManager.RemoveTask(taskID)
|
||||
}
|
||||
|
||||
// ValidateCron 验证 Cron 表达式
|
||||
@@ -494,10 +489,10 @@ func (es *ExecutorService) loadCronTasks() {
|
||||
go func(t models.Task) {
|
||||
// 延迟一点时间再触发,确保系统完全启动
|
||||
time.Sleep(3 * time.Second)
|
||||
logger.Infof("[Executor] 触发开机服务启动任务 #%d: %s", t.ID, t.Name)
|
||||
es.ExecuteTask(int(t.ID), nil)
|
||||
logger.Infof("[Executor] 触发开机服务启动任务 #%s: %s", t.ID, t.Name)
|
||||
es.ExecuteTask(t.ID, nil)
|
||||
}(task)
|
||||
} else if task.TriggerType == constant.TriggerTypeCron && task.Schedule != "" && (task.AgentID == nil || *task.AgentID == 0) {
|
||||
} else if task.TriggerType == constant.TriggerTypeCron && task.Schedule != "" && (task.AgentID == nil || *task.AgentID == "") {
|
||||
// 只调度本地任务(agent_id 为空或 0)的定时任务
|
||||
err := es.cronManager.AddTask(&task)
|
||||
if err != nil {
|
||||
@@ -519,11 +514,11 @@ func (es *ExecutorService) Reload() {
|
||||
}
|
||||
|
||||
// ExecuteTask executes a task by ID(同步执行,供 API 调用)
|
||||
func (es *ExecutorService) ExecuteTask(taskID int, extraEnvs []string) *executor.ExecutionResult {
|
||||
func (es *ExecutorService) ExecuteTask(taskID string, extraEnvs []string) *executor.ExecutionResult {
|
||||
task := es.taskService.GetTaskByID(taskID)
|
||||
if task == nil {
|
||||
return &executor.ExecutionResult{
|
||||
TaskID: fmt.Sprintf("%d", taskID),
|
||||
TaskID: taskID,
|
||||
Success: false,
|
||||
Error: "任务不存在",
|
||||
StartTime: time.Now(),
|
||||
@@ -532,9 +527,9 @@ func (es *ExecutorService) ExecuteTask(taskID int, extraEnvs []string) *executor
|
||||
}
|
||||
|
||||
// 1. 检查并发
|
||||
if err := es.CheckConcurrency(uint(taskID)); err != nil {
|
||||
if err := es.CheckConcurrency(taskID); err != nil {
|
||||
return &executor.ExecutionResult{
|
||||
TaskID: fmt.Sprintf("%d", taskID),
|
||||
TaskID: taskID,
|
||||
Success: false,
|
||||
Error: err.Error(), // 这里会返回 "任务正在运行中,拒绝并行执行"
|
||||
StartTime: time.Now(),
|
||||
@@ -548,7 +543,7 @@ func (es *ExecutorService) ExecuteTask(taskID int, extraEnvs []string) *executor
|
||||
}
|
||||
|
||||
req := &executor.ExecutionRequest{
|
||||
TaskID: fmt.Sprintf("%d", task.ID),
|
||||
TaskID: task.ID,
|
||||
Name: task.Name,
|
||||
Command: task.Command,
|
||||
WorkDir: task.WorkDir,
|
||||
@@ -562,7 +557,7 @@ func (es *ExecutorService) ExecuteTask(taskID int, extraEnvs []string) *executor
|
||||
es.scheduler.EnqueueOrExecute(req)
|
||||
|
||||
return &executor.ExecutionResult{
|
||||
TaskID: fmt.Sprintf("%d", task.ID),
|
||||
TaskID: task.ID,
|
||||
Success: true,
|
||||
Status: constant.TaskStatusQueued,
|
||||
StartTime: time.Now(),
|
||||
@@ -570,7 +565,7 @@ func (es *ExecutorService) ExecuteTask(taskID int, extraEnvs []string) *executor
|
||||
}
|
||||
|
||||
// StopTaskExecution stops a running task execution by LogID
|
||||
func (es *ExecutorService) StopTaskExecution(logID uint) error {
|
||||
func (es *ExecutorService) StopTaskExecution(logID string) error {
|
||||
var taskLog models.TaskLog
|
||||
if err := database.DB.First(&taskLog, logID).Error; err != nil {
|
||||
return fmt.Errorf("日志不存在")
|
||||
@@ -580,21 +575,21 @@ func (es *ExecutorService) StopTaskExecution(logID uint) error {
|
||||
return fmt.Errorf("任务已结束")
|
||||
}
|
||||
|
||||
task := es.taskService.GetTaskByID(int(taskLog.TaskID))
|
||||
task := es.taskService.GetTaskByID(taskLog.TaskID)
|
||||
if task == nil {
|
||||
return fmt.Errorf("任务不存在")
|
||||
}
|
||||
|
||||
// 远程任务:发送停止指令到 Agent
|
||||
if task.AgentID != nil && *task.AgentID > 0 {
|
||||
logger.Infof("[Executor] 请求停止远程任务 #%d (Agent #%d, LogID: %d)", task.ID, *task.AgentID, logID)
|
||||
if task.AgentID != nil && *task.AgentID != "" {
|
||||
logger.Infof("[Executor] 请求停止远程任务 #%s (Agent #%s, LogID: %s)", task.ID, *task.AgentID, logID)
|
||||
return es.agentWSManager.SendToAgent(*task.AgentID, constant.WSTypeStop, map[string]interface{}{
|
||||
"log_id": logID,
|
||||
})
|
||||
}
|
||||
|
||||
// 本地任务:直接停止调度器中的执行实例
|
||||
logger.Infof("[Executor] 请求停止本地任务 #%d (LogID: %d)", task.ID, logID)
|
||||
logger.Infof("[Executor] 请求停止本地任务 #%s (LogID: %s)", task.ID, logID)
|
||||
if es.scheduler.StopLog(logID) {
|
||||
return nil
|
||||
}
|
||||
@@ -658,7 +653,7 @@ func (es *ExecutorService) UpdateResult(res executor.ExecutionResult) {
|
||||
|
||||
// 查找是否已存在(通过 LogID)
|
||||
for i := range es.results {
|
||||
if es.results[i].LogID == res.LogID && res.LogID != 0 {
|
||||
if es.results[i].LogID == res.LogID && res.LogID != "" {
|
||||
es.results[i] = res
|
||||
return
|
||||
}
|
||||
@@ -703,9 +698,9 @@ func (es *ExecutorService) CleanupRunningTasks() error {
|
||||
}
|
||||
|
||||
// CheckConcurrency 检查任务并发限制(只读检查)
|
||||
func (es *ExecutorService) CheckConcurrency(taskID uint) error {
|
||||
func (es *ExecutorService) CheckConcurrency(taskID string) error {
|
||||
var task models.Task
|
||||
if err := database.DB.Select("config, running_go").First(&task, taskID).Error; err != nil {
|
||||
if err := database.DB.Select("config, running_go").Where("id = ?", taskID).First(&task).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var goids []int64
|
||||
@@ -725,11 +720,11 @@ func (es *ExecutorService) CheckConcurrency(taskID uint) error {
|
||||
}
|
||||
|
||||
// AddRunningGo 添加当前 goroutine ID 到任务的 running_go 字段
|
||||
func (es *ExecutorService) AddRunningGo(taskID uint) (int64, error) {
|
||||
func (es *ExecutorService) AddRunningGo(taskID string) (int64, error) {
|
||||
goid := utils.GetGoroutineID()
|
||||
err := database.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var task models.Task
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&task, taskID).Error; err != nil {
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("id = ?", taskID).First(&task).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var goids []int64
|
||||
@@ -756,10 +751,10 @@ func (es *ExecutorService) AddRunningGo(taskID uint) (int64, error) {
|
||||
}
|
||||
|
||||
// RemoveRunningGo 从任务的 running_go 字段移除指定 goroutine ID
|
||||
func (es *ExecutorService) RemoveRunningGo(taskID uint, goid int64) {
|
||||
func (es *ExecutorService) RemoveRunningGo(taskID string, goid int64) {
|
||||
database.DB.Transaction(func(tx *gorm.DB) error {
|
||||
var task models.Task
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&task, taskID).Error; err != nil {
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("id = ?", taskID).First(&task).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
var goids []int64
|
||||
@@ -778,17 +773,17 @@ func (es *ExecutorService) RemoveRunningGo(taskID uint, goid int64) {
|
||||
}
|
||||
|
||||
// ExecuteRemoteForScheduler 供 Scheduler 调用,执行远程任务并等待结果
|
||||
func (es *ExecutorService) ExecuteRemoteForScheduler(task *models.Task, logID uint) (*executor.Result, error) {
|
||||
func (es *ExecutorService) ExecuteRemoteForScheduler(task *models.Task, logID string) (*executor.Result, error) {
|
||||
agentID := *task.AgentID
|
||||
logger.Infof("[Executor] 远程执行任务 #%d: %s (Agent #%d, LogID: %d)", task.ID, task.Name, agentID, logID)
|
||||
logger.Infof("[Executor] 远程执行任务 #%s: %s (Agent #%s, LogID: %s)", task.ID, task.Name, agentID, logID)
|
||||
|
||||
// 1. 检查 Agent 状态
|
||||
var agent models.Agent
|
||||
if err := database.DB.First(&agent, agentID).Error; err != nil {
|
||||
return nil, fmt.Errorf("Agent #%d 不存在", agentID)
|
||||
if err := database.DB.Where("id = ?", agentID).First(&agent).Error; err != nil {
|
||||
return nil, fmt.Errorf("Agent #%s 不存在", agentID)
|
||||
}
|
||||
if !agent.Enabled {
|
||||
return nil, fmt.Errorf("Agent #%d 已禁用", agentID)
|
||||
return nil, fmt.Errorf("Agent #%s 已禁用", agentID)
|
||||
}
|
||||
if es.agentWSManager == nil {
|
||||
return nil, fmt.Errorf("AgentWSManager 未初始化")
|
||||
|
||||
@@ -13,7 +13,7 @@ import (
|
||||
|
||||
// SendStatsService 接口定义(避免循环依赖)
|
||||
type SendStatsService interface {
|
||||
IncrementStats(taskID uint, status string) error
|
||||
IncrementStats(taskID string, status string) error
|
||||
}
|
||||
|
||||
// TaskLogService 任务日志服务
|
||||
@@ -35,9 +35,10 @@ type CleanConfig struct {
|
||||
}
|
||||
|
||||
// CreateEmptyLog 创建一个空的日志记录(任务开始时调用)
|
||||
func (s *TaskLogService) CreateEmptyLog(taskID uint, command string) (*models.TaskLog, error) {
|
||||
func (s *TaskLogService) CreateEmptyLog(taskID string, command string) (*models.TaskLog, error) {
|
||||
startTime := models.Now()
|
||||
taskLog := &models.TaskLog{
|
||||
ID: utils.GenerateID(),
|
||||
TaskID: taskID,
|
||||
Command: command,
|
||||
Status: "running",
|
||||
@@ -52,9 +53,10 @@ func (s *TaskLogService) CreateEmptyLog(taskID uint, command string) (*models.Ta
|
||||
// SaveTaskLog 保存或更新任务日志
|
||||
func (s *TaskLogService) SaveTaskLog(taskLog *models.TaskLog) error {
|
||||
var err error
|
||||
if taskLog.ID > 0 {
|
||||
err = database.DB.Model(taskLog).Updates(taskLog).Error
|
||||
if taskLog.ID != "" {
|
||||
err = database.DB.Model(taskLog).Where("id = ?", taskLog.ID).Updates(taskLog).Error
|
||||
} else {
|
||||
taskLog.ID = utils.GenerateID()
|
||||
err = database.DB.Create(taskLog).Error
|
||||
}
|
||||
|
||||
@@ -69,12 +71,12 @@ func (s *TaskLogService) SaveTaskLog(taskLog *models.TaskLog) error {
|
||||
}
|
||||
|
||||
// UpdateTaskDuration 更新任务耗时(心跳)
|
||||
func (s *TaskLogService) UpdateTaskDuration(logID uint, duration int64) error {
|
||||
func (s *TaskLogService) UpdateTaskDuration(logID string, duration int64) error {
|
||||
return database.DB.Model(&models.TaskLog{}).Where("id = ?", logID).Update("duration", duration).Error
|
||||
}
|
||||
|
||||
// UpdateTaskStats 更新任务统计
|
||||
func (s *TaskLogService) UpdateTaskStats(taskID uint, status string) {
|
||||
func (s *TaskLogService) UpdateTaskStats(taskID string, status string) {
|
||||
if s.sendStatsService == nil {
|
||||
logger.Error("[TaskLog] SendStatsService 未初始化")
|
||||
return
|
||||
@@ -87,9 +89,9 @@ func (s *TaskLogService) UpdateTaskStats(taskID uint, status string) {
|
||||
}
|
||||
|
||||
// CleanTaskLogs 清理任务日志
|
||||
func (s *TaskLogService) CleanTaskLogs(taskID uint) {
|
||||
func (s *TaskLogService) CleanTaskLogs(taskID string) {
|
||||
var task models.Task
|
||||
if err := database.DB.First(&task, taskID).Error; err != nil {
|
||||
if err := database.DB.Where("id = ?", taskID).First(&task).Error; err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -123,7 +125,7 @@ func (s *TaskLogService) CleanTaskLogs(taskID uint) {
|
||||
}
|
||||
|
||||
if deleted > 0 {
|
||||
logger.Infof("[TaskLog] 清理任务 #%d 的 %d 条日志", taskID, deleted)
|
||||
logger.Infof("[TaskLog] 清理任务 #%s 的 %d 条日志", taskID, deleted)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -153,6 +155,7 @@ func (s *TaskLogService) CreateTaskLogFromAgentResult(result *models.AgentTaskRe
|
||||
}
|
||||
|
||||
taskLog := &models.TaskLog{
|
||||
ID: utils.GenerateID(),
|
||||
TaskID: result.TaskID,
|
||||
AgentID: &result.AgentID,
|
||||
Command: result.Command,
|
||||
@@ -177,7 +180,7 @@ func (s *TaskLogService) CreateTaskLogFromAgentResult(result *models.AgentTaskRe
|
||||
}
|
||||
|
||||
// CreateTaskLogFromLocalExecution 从本地执行结果创建任务日志
|
||||
func (s *TaskLogService) CreateTaskLogFromLocalExecution(taskID uint, command, output, systemErr, status string, duration int64, exitCode int, start, end time.Time, isCompressed bool) (*models.TaskLog, error) {
|
||||
func (s *TaskLogService) CreateTaskLogFromLocalExecution(taskID string, command, output, systemErr, status string, duration int64, exitCode int, start, end time.Time, isCompressed bool) (*models.TaskLog, error) {
|
||||
var compressed string
|
||||
var err error
|
||||
|
||||
@@ -196,6 +199,7 @@ func (s *TaskLogService) CreateTaskLogFromLocalExecution(taskID uint, command, o
|
||||
endTime := models.LocalTime(end)
|
||||
|
||||
taskLog := &models.TaskLog{
|
||||
ID: utils.GenerateID(),
|
||||
TaskID: taskID,
|
||||
Command: command,
|
||||
Output: compressed,
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"github.com/engigu/baihu-panel/internal/constant"
|
||||
"github.com/engigu/baihu-panel/internal/database"
|
||||
"github.com/engigu/baihu-panel/internal/models"
|
||||
"github.com/engigu/baihu-panel/internal/utils"
|
||||
)
|
||||
|
||||
type TaskService struct{}
|
||||
@@ -12,7 +13,7 @@ func NewTaskService() *TaskService {
|
||||
return &TaskService{}
|
||||
}
|
||||
|
||||
func (ts *TaskService) CreateTask(name, command, schedule string, timeout int, workDir, cleanConfig, envs, taskType, config string, agentID *uint, languages []map[string]string, triggerType string, tags string, retryCount int, retryInterval int, randomRange int) *models.Task {
|
||||
func (ts *TaskService) CreateTask(name, command, schedule string, timeout int, workDir, cleanConfig, envs, taskType, config string, agentID *string, languages []map[string]string, triggerType string, tags string, retryCount int, retryInterval int, randomRange int) *models.Task {
|
||||
if taskType == "" {
|
||||
taskType = "task"
|
||||
}
|
||||
@@ -20,6 +21,7 @@ func (ts *TaskService) CreateTask(name, command, schedule string, timeout int, w
|
||||
triggerType = constant.TriggerTypeCron
|
||||
}
|
||||
task := &models.Task{
|
||||
ID: utils.GenerateID(),
|
||||
Name: name,
|
||||
Command: command,
|
||||
Tags: tags,
|
||||
@@ -52,7 +54,7 @@ func (ts *TaskService) GetTasks() []models.Task {
|
||||
}
|
||||
|
||||
// GetTasksWithPagination 分页获取任务列表
|
||||
func (ts *TaskService) GetTasksWithPagination(page, pageSize int, name string, agentID *uint, tags string, taskType string) ([]models.Task, int64) {
|
||||
func (ts *TaskService) GetTasksWithPagination(page, pageSize int, name string, agentID *string, tags string, taskType string) ([]models.Task, int64) {
|
||||
var tasks []models.Task
|
||||
var total int64
|
||||
|
||||
@@ -76,17 +78,17 @@ func (ts *TaskService) GetTasksWithPagination(page, pageSize int, name string, a
|
||||
return tasks, total
|
||||
}
|
||||
|
||||
func (ts *TaskService) GetTaskByID(id int) *models.Task {
|
||||
func (ts *TaskService) GetTaskByID(id string) *models.Task {
|
||||
var task models.Task
|
||||
if err := database.DB.First(&task, id).Error; err != nil {
|
||||
if err := database.DB.Where("id = ?", id).First(&task).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
return &task
|
||||
}
|
||||
|
||||
func (ts *TaskService) UpdateTask(id int, name, command, schedule string, timeout int, workDir, cleanConfig, envs string, enabled bool, taskType, config string, agentID *uint, languages []map[string]string, triggerType string, tags string, retryCount int, retryInterval int, randomRange int) *models.Task {
|
||||
func (ts *TaskService) UpdateTask(id string, name, command, schedule string, timeout int, workDir, cleanConfig, envs string, enabled bool, taskType, config string, agentID *string, languages []map[string]string, triggerType string, tags string, retryCount int, retryInterval int, randomRange int) *models.Task {
|
||||
var task models.Task
|
||||
if err := database.DB.First(&task, id).Error; err != nil {
|
||||
if err := database.DB.Where("id = ?", id).First(&task).Error; err != nil {
|
||||
return nil
|
||||
}
|
||||
task.Name = name
|
||||
@@ -117,7 +119,7 @@ func (ts *TaskService) UpdateTask(id int, name, command, schedule string, timeou
|
||||
return &task
|
||||
}
|
||||
|
||||
func (ts *TaskService) DeleteTask(id int) bool {
|
||||
result := database.DB.Delete(&models.Task{}, id)
|
||||
func (ts *TaskService) DeleteTask(id string) bool {
|
||||
result := database.DB.Where("id = ?", id).Delete(&models.Task{})
|
||||
return result.RowsAffected > 0
|
||||
}
|
||||
|
||||
@@ -16,13 +16,13 @@ import (
|
||||
var (
|
||||
// globalTinyLogManager 跟踪所有活跃的 TinyLog 实例
|
||||
globalTinyLogManager = &TinyLogManager{
|
||||
logs: make(map[uint]*TinyLog),
|
||||
logs: make(map[string]*TinyLog),
|
||||
}
|
||||
)
|
||||
|
||||
type TinyLogManager struct {
|
||||
mu sync.RWMutex
|
||||
logs map[uint]*TinyLog
|
||||
logs map[string]*TinyLog
|
||||
}
|
||||
|
||||
func (m *TinyLogManager) Register(log *TinyLog) {
|
||||
@@ -31,26 +31,26 @@ func (m *TinyLogManager) Register(log *TinyLog) {
|
||||
m.logs[log.LogID] = log
|
||||
}
|
||||
|
||||
func (m *TinyLogManager) Unregister(logID uint) {
|
||||
func (m *TinyLogManager) Unregister(logID string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.logs, logID)
|
||||
}
|
||||
|
||||
func (m *TinyLogManager) Get(logID uint) *TinyLog {
|
||||
func (m *TinyLogManager) Get(logID string) *TinyLog {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.logs[logID]
|
||||
}
|
||||
|
||||
// GetActiveLog 通过 ID 获取活跃的 TinyLog 实例
|
||||
func GetActiveLog(logID uint) *TinyLog {
|
||||
func GetActiveLog(logID string) *TinyLog {
|
||||
return globalTinyLogManager.Get(logID)
|
||||
}
|
||||
|
||||
// TinyLog 是一个高性能、低内存占用的日志收集器
|
||||
type TinyLog struct {
|
||||
LogID uint
|
||||
LogID string
|
||||
mu sync.RWMutex
|
||||
file *os.File
|
||||
path string
|
||||
@@ -61,7 +61,7 @@ type TinyLog struct {
|
||||
}
|
||||
|
||||
// NewTinyLog 创建一个新的 TinyLog 实例(基于临时文件存储)并注册它
|
||||
func NewTinyLog(logID uint) (*TinyLog, error) {
|
||||
func NewTinyLog(logID string) (*TinyLog, error) {
|
||||
f, err := os.CreateTemp("", "task_log_*.log")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"github.com/engigu/baihu-panel/internal/constant"
|
||||
"github.com/engigu/baihu-panel/internal/database"
|
||||
"github.com/engigu/baihu-panel/internal/models"
|
||||
"github.com/engigu/baihu-panel/internal/utils"
|
||||
)
|
||||
|
||||
type UserService struct{}
|
||||
@@ -22,6 +23,7 @@ func (us *UserService) hashPassword(password string) string {
|
||||
|
||||
func (us *UserService) CreateUser(username, password, email, role string) *models.User {
|
||||
user := &models.User{
|
||||
ID: utils.GenerateID(),
|
||||
Username: username,
|
||||
Password: us.hashPassword(password),
|
||||
Email: email,
|
||||
@@ -59,6 +61,6 @@ func (us *UserService) AuthenticateUser(username, password string) bool {
|
||||
return us.ValidatePassword(user, password)
|
||||
}
|
||||
|
||||
func (us *UserService) UpdatePassword(userID uint, newPassword string) error {
|
||||
func (us *UserService) UpdatePassword(userID string, newPassword string) error {
|
||||
return database.DB.Model(&models.User{}).Where("id = ?", userID).Update("password", us.hashPassword(newPassword)).Error
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user