diff --git a/docker/docker-entrypoint.sh b/docker/docker-entrypoint.sh index f722e7e..a72179e 100644 --- a/docker/docker-entrypoint.sh +++ b/docker/docker-entrypoint.sh @@ -77,6 +77,7 @@ log " - node: $(node --version 2>&1 | head -n 1 || echo "not found")" # 延迟获取 NODE_PATH,避免同步阻塞启动 log "Checking npm..." log " - npm: $(npm --version 2>&1 | head -n 1 || echo "not found")" + export NODE_PATH=$(npm root -g 2>/dev/null || echo "") log " - node_path: $NODE_PATH" diff --git a/internal/controllers/log_controller.go b/internal/controllers/log_controller.go index 4d5ab3c..031b11d 100644 --- a/internal/controllers/log_controller.go +++ b/internal/controllers/log_controller.go @@ -118,7 +118,8 @@ func (lc *LogController) GetLogDetail(c *gin.Context) { } var log models.TaskLog - if err := database.DB.Where("id = ?", id).First(&log).Error; err != nil { + res := database.DB.Where("id = ?", id).Limit(1).Find(&log) + if res.Error != nil || res.RowsAffected == 0 { utils.NotFound(c, "日志不存在") return } diff --git a/internal/controllers/log_ws_controller.go b/internal/controllers/log_ws_controller.go index 9c09887..4c77135 100644 --- a/internal/controllers/log_ws_controller.go +++ b/internal/controllers/log_ws_controller.go @@ -34,7 +34,8 @@ func (lc *LogWSController) StreamLog(c *gin.Context) { // 1. 检查数据库中是否已结束 var taskLog models.TaskLog - if err := database.DB.Where("id = ?", logID).First(&taskLog).Error; err == nil { + res := database.DB.Where("id = ?", logID).Limit(1).Find(&taskLog) + if res.Error == nil && res.RowsAffected > 0 { if taskLog.Status != "running" { // 已结束,读取库内日志 content, err := utils.DecompressFromBase64(string(taskLog.Output)) @@ -74,7 +75,8 @@ func (lc *LogWSController) StreamLog(c *gin.Context) { if !ok { // 任务结束,尝试刷新最后一次库内完整内容 var finalLog models.TaskLog - if err := database.DB.Where("id = ?", logID).First(&finalLog).Error; err == nil { + res := database.DB.Where("id = ?", logID).Limit(1).Find(&finalLog) + if res.Error == nil && res.RowsAffected > 0 { content, _ := utils.DecompressFromBase64(string(finalLog.Output)) if content != "" { conn.WriteMessage(websocket.TextMessage, []byte("\n--- 任务已结束 ---\n")) diff --git a/internal/controllers/settings_controller.go b/internal/controllers/settings_controller.go index eae020b..e24d92d 100644 --- a/internal/controllers/settings_controller.go +++ b/internal/controllers/settings_controller.go @@ -65,7 +65,8 @@ func (sc *SettingsController) ChangePassword(c *gin.Context) { userID := c.GetString("userID") var user *models.User - if err := database.DB.Where("id = ?", userID).First(&user).Error; err != nil { + res := database.DB.Where("id = ?", userID).Limit(1).Find(&user) + if res.Error != nil || res.RowsAffected == 0 { utils.NotFound(c, "用户不存在") return } diff --git a/internal/middleware/auth.go b/internal/middleware/auth.go index c2e2060..bdd1129 100644 --- a/internal/middleware/auth.go +++ b/internal/middleware/auth.go @@ -47,7 +47,8 @@ func AuthRequired() gin.HandlerFunc { // 安全增强:校验数据库中该用户的 ID 是否与 Token 一致,并验证 TokenVersion var user models.User - if err := database.DB.Where("username = ?", username).First(&user).Error; err != nil || user.ID != userID || user.TokenVersion != tokenVersion { + res := database.DB.Where("username = ?", username).Limit(1).Find(&user) + if res.Error != nil || res.RowsAffected == 0 || user.ID != userID || user.TokenVersion != tokenVersion { utils.Unauthorized(c, "会话失效,请重新登录") ClearAuthCookie(c) c.Abort() @@ -150,7 +151,8 @@ func checkOpenapiToken(c *gin.Context, settingsSvc *services.SettingsService) bo // 模拟 Admin 角色 var adminUser models.User - if err := database.DB.Where("role = ?", "admin").First(&adminUser).Error; err != nil { + res := database.DB.Where("role = ?", "admin").Limit(1).Find(&adminUser) + if res.Error != nil || res.RowsAffected == 0 { utils.Unauthorized(c, "未找到管理员账户,OpenAPI Token 校验失败") c.Abort() return true diff --git a/internal/services/agent_service.go b/internal/services/agent_service.go index d59f24b..9d0611a 100644 --- a/internal/services/agent_service.go +++ b/internal/services/agent_service.go @@ -80,7 +80,8 @@ func (s *AgentService) DeleteToken(id string) error { // ValidateToken 验证令牌 func (s *AgentService) ValidateToken(token string) (*models.AgentToken, error) { var agentToken models.AgentToken - if err := database.DB.Where("token = ?", token).First(&agentToken).Error; err != nil { + res := database.DB.Where("token = ?", token).Limit(1).Find(&agentToken) + if res.Error != nil || res.RowsAffected == 0 { return nil, &ServiceError{Message: "无效的令牌"} } @@ -120,7 +121,8 @@ func (s *AgentService) RegisterByToken(token string, machineID string, ip string // 如果提供了 machine_id,先检查是否已存在 if machineID != "" { var existing models.Agent - if err := database.DB.Where("machine_id = ?", machineID).First(&existing).Error; err == nil { + res := database.DB.Where("machine_id = ?", machineID).Limit(1).Find(&existing) + if res.Error == nil && res.RowsAffected > 0 { // 已存在,更新 token 和状态,复用已有 Agent now := models.LocalTime(time.Now()) database.DB.Model(&existing).Updates(map[string]interface{}{ @@ -171,7 +173,8 @@ func (s *AgentService) Register(req *models.AgentRegisterRequest, ip string) (*m // 检查是否已存在同名 Agent var existing models.Agent - if err := database.DB.Where("name = ?", req.Name).First(&existing).Error; err == nil { + res := database.DB.Where("name = ?", req.Name).Limit(1).Find(&existing) + if res.Error == nil && res.RowsAffected > 0 { return nil, "", &ServiceError{Message: "Agent 名称已存在"} } @@ -223,7 +226,8 @@ func (s *AgentService) Delete(id string) error { // GetByID 根据 ID 获取 Agent func (s *AgentService) GetByID(id string) *models.Agent { var agent models.Agent - if err := database.DB.Where("id = ?", id).First(&agent).Error; err != nil { + res := database.DB.Where("id = ?", id).Limit(1).Find(&agent) + if res.Error != nil || res.RowsAffected == 0 { return nil } return &agent @@ -232,7 +236,8 @@ func (s *AgentService) GetByID(id string) *models.Agent { // GetByToken 根据 Token 获取 Agent func (s *AgentService) GetByToken(token string) *models.Agent { var agent models.Agent - if err := database.DB.Where("token = ?", token).First(&agent).Error; err != nil { + res := database.DB.Where("token = ?", token).Limit(1).Find(&agent) + if res.Error != nil || res.RowsAffected == 0 { return nil } return &agent @@ -241,7 +246,8 @@ func (s *AgentService) GetByToken(token string) *models.Agent { // GetByMachineID 根据 MachineID 获取 Agent func (s *AgentService) GetByMachineID(machineID string) *models.Agent { var agent models.Agent - if err := database.DB.Where("machine_id = ?", machineID).First(&agent).Error; err != nil { + res := database.DB.Where("machine_id = ?", machineID).Limit(1).Find(&agent) + if res.Error != nil || res.RowsAffected == 0 { return nil } return &agent diff --git a/internal/services/backup_service.go b/internal/services/backup_service.go index 1ae6823..aef7f7a 100644 --- a/internal/services/backup_service.go +++ b/internal/services/backup_service.go @@ -417,7 +417,8 @@ func (s *BackupService) addDirToZip(zipWriter *zip.Writer, srcDir, prefix string func (s *BackupService) GetBackupFile() string { var setting models.Setting - if err := database.DB.Where("section = ? AND `key` = ?", BackupSection, BackupFileKey).First(&setting).Error; err != nil { + res := database.DB.Where("section = ? AND `key` = ?", BackupSection, BackupFileKey).Limit(1).Find(&setting) + if res.Error != nil || res.RowsAffected == 0 { return "" } return string(setting.Value) diff --git a/internal/services/dependency_service.go b/internal/services/dependency_service.go index 5309a9b..dc9ef77 100644 --- a/internal/services/dependency_service.go +++ b/internal/services/dependency_service.go @@ -33,8 +33,8 @@ func (s *DependencyService) List(language, langVersion string) ([]models.Depende func (s *DependencyService) Create(dep *models.Dependency) error { // 检查是否已存在(名称、版本、语言及版本必须完全匹配) var existing models.Dependency - err := database.DB.Where("name = ? AND version = ? AND language = ? AND lang_version = ?", dep.Name, dep.Version, dep.Language, dep.LangVersion).First(&existing).Error - if err == nil { + res := database.DB.Where("name = ? AND version = ? AND language = ? AND lang_version = ?", dep.Name, dep.Version, dep.Language, dep.LangVersion).Limit(1).Find(&existing) + if res.Error == nil && res.RowsAffected > 0 { // 如果已存在,更新 ID 并执行更新 dep.ID = existing.ID return database.DB.Model(&existing).Updates(dep).Error diff --git a/internal/services/env_service.go b/internal/services/env_service.go index 646a8dd..7f8ff71 100644 --- a/internal/services/env_service.go +++ b/internal/services/env_service.go @@ -60,7 +60,8 @@ func (es *EnvService) GetEnvVarsWithPagination(userID string, name string, page, func (es *EnvService) GetEnvVarByID(id string) *models.EnvironmentVariable { var env models.EnvironmentVariable - if err := database.DB.Where("id = ?", id).First(&env).Error; err != nil { + res := database.DB.Where("id = ?", id).Limit(1).Find(&env) + if res.Error != nil || res.RowsAffected == 0 { return nil } return &env @@ -68,7 +69,8 @@ func (es *EnvService) GetEnvVarByID(id string) *models.EnvironmentVariable { func (es *EnvService) UpdateEnvVar(id string, name, value, remark string, hidden, enabled bool) *models.EnvironmentVariable { var env models.EnvironmentVariable - if err := database.DB.Where("id = ?", id).First(&env).Error; err != nil { + res := database.DB.Where("id = ?", id).Limit(1).Find(&env) + if res.Error != nil || res.RowsAffected == 0 { return nil } updates := map[string]interface{}{ diff --git a/internal/services/migration_v3.go b/internal/services/migration_v3.go index b316631..8680ff6 100644 --- a/internal/services/migration_v3.go +++ b/internal/services/migration_v3.go @@ -81,8 +81,8 @@ func RunMigrationV3() error { // 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" { + res := db.Where("section = ? AND `key` = ?", "system", "migration_v3_success").Limit(1).Find(&migrationFlag) + if res.Error == nil && res.RowsAffected > 0 && migrationFlag.Value == "true" { // 如果已经是字符串 ID 模式,双重确认 return nil } @@ -150,8 +150,8 @@ func markMigrationSuccess(db *gorm.DB) error { return nil } var flag models.Setting - err := db.Where("section = ? AND `key` = ?", "system", "migration_v3_success").First(&flag).Error - if err != nil { + res := db.Where("section = ? AND `key` = ?", "system", "migration_v3_success").Limit(1).Find(&flag) + if res.Error != nil || res.RowsAffected == 0 { // 创建或更新 flag = models.Setting{ ID: utils.GenerateID(), diff --git a/internal/services/mise_service.go b/internal/services/mise_service.go index da59614..1d094cf 100644 --- a/internal/services/mise_service.go +++ b/internal/services/mise_service.go @@ -271,7 +271,9 @@ func (s *MiseService) syncToDB(languages []MiseLanguage) { for _, lang := range languages { var model models.Language // 以 plugin 和 version 作为联合唯一标识(业务逻辑上) - err := db.Where("plugin = ? AND version = ?", lang.Plugin, lang.Version).First(&model).Error + res := db.Where("plugin = ? AND version = ?", lang.Plugin, lang.Version).Limit(1).Find(&model) + err := res.Error + rowsAffected := res.RowsAffected sourceStr := "" if lang.Source.Path != "" { @@ -289,8 +291,8 @@ func (s *MiseService) syncToDB(languages []MiseLanguage) { } } - if err != nil { - // 如果不存在,则创建 + if err == nil && rowsAffected == 0 { + // 如果不存在,则创建 newLang := models.Language{ ID: utils.GenerateID(), Plugin: lang.Plugin, diff --git a/internal/services/notification_service.go b/internal/services/notification_service.go index bbe0c9e..7de53f8 100644 --- a/internal/services/notification_service.go +++ b/internal/services/notification_service.go @@ -152,9 +152,9 @@ func (s *NotificationService) SaveBinding(binding *models.NotifyBinding) error { if binding.ID == "" { // 检查是否已经存在相同的绑定(避免重复点击导致多个记录) var existing models.NotifyBinding - err := database.DB.Where("type = ? AND event = ? AND way_id = ? AND data_id = ?", - binding.Type, binding.Event, binding.WayID, binding.DataID).First(&existing).Error - if err == nil { + res := database.DB.Where("type = ? AND event = ? AND way_id = ? AND data_id = ?", + binding.Type, binding.Event, binding.WayID, binding.DataID).Limit(1).Find(&existing) + if res.Error == nil && res.RowsAffected > 0 { // 如果已存在且未删除,直接返回(或者更新它) *binding = existing return nil @@ -246,7 +246,8 @@ func (s *NotificationService) SendByChannelID(channelID string, msg *NotifyMessa defer s.mu.RUnlock() var notifyWay models.NotifyWay - if err := database.DB.Where("id = ?", channelID).First(¬ifyWay).Error; err != nil { + res := database.DB.Where("id = ?", channelID).Limit(1).Find(¬ifyWay) + if res.Error != nil || res.RowsAffected == 0 { return &NotifyResult{Success: false, Error: "渠道不存在"} } diff --git a/internal/services/script_service.go b/internal/services/script_service.go index ead4994..77ce604 100644 --- a/internal/services/script_service.go +++ b/internal/services/script_service.go @@ -31,7 +31,8 @@ func (ss *ScriptService) GetScriptsByUserID(userID string) []models.Script { func (ss *ScriptService) GetScriptByID(id string) *models.Script { var script models.Script - if err := database.DB.Where("id = ?", id).First(&script).Error; err != nil { + res := database.DB.Where("id = ?", id).Limit(1).Find(&script) + if res.Error != nil || res.RowsAffected == 0 { return nil } return &script @@ -39,7 +40,8 @@ func (ss *ScriptService) GetScriptByID(id string) *models.Script { func (ss *ScriptService) UpdateScript(id string, name, content string) *models.Script { var script models.Script - if err := database.DB.Where("id = ?", id).First(&script).Error; err != nil { + res := database.DB.Where("id = ?", id).Limit(1).Find(&script) + if res.Error != nil || res.RowsAffected == 0 { return nil } script.Name = name diff --git a/internal/services/send_stats_service.go b/internal/services/send_stats_service.go index e10abba..c133c06 100644 --- a/internal/services/send_stats_service.go +++ b/internal/services/send_stats_service.go @@ -20,9 +20,8 @@ func (s *SendStatsService) IncrementStats(taskID string, status string) error { day := systime.FormatDate(time.Now()) var stats models.SendStats - result := database.DB.Where("task_id = ? AND day = ? AND status = ?", taskID, day, status).First(&stats) - - if result.Error != nil { + res := database.DB.Where("task_id = ? AND day = ? AND status = ?", taskID, day, status).Limit(1).Find(&stats) + if res.Error != nil || res.RowsAffected == 0 { // 不存在则创建 stats = models.SendStats{ ID: utils.GenerateID(), diff --git a/internal/services/settings_service.go b/internal/services/settings_service.go index 3f3362a..f2045d2 100644 --- a/internal/services/settings_service.go +++ b/internal/services/settings_service.go @@ -121,7 +121,8 @@ func (s *SettingsService) Get(section, key string) string { return cache.GetSiteCache(key) } var setting models.Setting - if err := database.DB.Where("section = ? AND `key` = ?", section, key).First(&setting).Error; err != nil { + res := database.DB.Where("section = ? AND `key` = ?", section, key).Limit(1).Find(&setting) + if res.Error != nil || res.RowsAffected == 0 { if def, ok := constant.DefaultSettings[section][key]; ok { return def } @@ -133,7 +134,8 @@ func (s *SettingsService) Get(section, key string) string { // Set 设置单个值 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 { + res := database.DB.Where("section = ? AND `key` = ?", section, key).Limit(1).Find(&setting) + if res.Error != nil || res.RowsAffected == 0 { return database.DB.Create(&models.Setting{ ID: utils.GenerateID(), Section: section, diff --git a/internal/services/tasks/executor_service.go b/internal/services/tasks/executor_service.go index 417f079..9da2463 100644 --- a/internal/services/tasks/executor_service.go +++ b/internal/services/tasks/executor_service.go @@ -614,7 +614,8 @@ func (es *ExecutorService) ExecuteTask(taskID string, extraEnvs []string) *execu // StopTaskExecution stops a running task execution by LogID func (es *ExecutorService) StopTaskExecution(logID string) error { var taskLog models.TaskLog - if err := database.DB.Where("id = ?", logID).First(&taskLog).Error; err != nil { + res := database.DB.Where("id = ?", logID).Limit(1).Find(&taskLog) + if res.Error != nil || res.RowsAffected == 0 { return fmt.Errorf("日志不存在") } @@ -747,7 +748,8 @@ func (es *ExecutorService) CleanupRunningTasks() error { // CheckConcurrency 检查任务并发限制(只读检查) func (es *ExecutorService) CheckConcurrency(taskID string) error { var task models.Task - if err := database.DB.Select("config, running_go").Where("id = ?", taskID).First(&task).Error; err != nil { + res := database.DB.Select("config, running_go").Where("id = ?", taskID).Limit(1).Find(&task) + if res.Error != nil || res.RowsAffected == 0 { return err } var goids []int64 @@ -773,7 +775,8 @@ func (es *ExecutorService) AddRunningGo(taskID string) (int64, error) { for attempt := 0; attempt < 3; attempt++ { lastErr = database.DB.Transaction(func(tx *gorm.DB) error { var task models.Task - if err := tx.Where("id = ?", taskID).First(&task).Error; err != nil { + res := tx.Where("id = ?", taskID).Limit(1).Find(&task) + if res.Error != nil || res.RowsAffected == 0 { return err } var goids []int64 @@ -814,7 +817,8 @@ func (es *ExecutorService) RemoveRunningGo(taskID string, goid int64) { for attempt := 0; attempt < 3; attempt++ { err := database.DB.Transaction(func(tx *gorm.DB) error { var task models.Task - if err := tx.Where("id = ?", taskID).First(&task).Error; err != nil { + res := tx.Where("id = ?", taskID).Limit(1).Find(&task) + if res.Error != nil || res.RowsAffected == 0 { return err } var goids []int64 @@ -844,7 +848,8 @@ func (es *ExecutorService) ExecuteRemoteForScheduler(task *models.Task, logID st // 1. 检查 Agent 状态 var agent models.Agent - if err := database.DB.Where("id = ?", agentID).First(&agent).Error; err != nil { + res := database.DB.Where("id = ?", agentID).Limit(1).Find(&agent) + if res.Error != nil || res.RowsAffected == 0 { return nil, fmt.Errorf("Agent #%s 不存在", agentID) } if !agent.Enabled { diff --git a/internal/services/tasks/ql_repo_parser.go b/internal/services/tasks/ql_repo_parser.go index 2ed71c0..3aacfcf 100644 --- a/internal/services/tasks/ql_repo_parser.go +++ b/internal/services/tasks/ql_repo_parser.go @@ -38,7 +38,8 @@ func ParseRepoScriptsAndAddCron(es *ExecutorService, taskID string, logWriter io } var repoTask models.Task - if err := database.DB.Where("id = ?", taskID).First(&repoTask).Error; err != nil { + res := database.DB.Where("id = ?", taskID).Limit(1).Find(&repoTask) + if res.Error != nil || res.RowsAffected == 0 { return } diff --git a/internal/services/tasks/task_log_service.go b/internal/services/tasks/task_log_service.go index 26e1030..8a7293a 100644 --- a/internal/services/tasks/task_log_service.go +++ b/internal/services/tasks/task_log_service.go @@ -95,7 +95,8 @@ func (s *TaskLogService) UpdateTaskStats(taskID string, status string) { // CleanTaskLogs 清理任务日志 func (s *TaskLogService) CleanTaskLogs(taskID string) { var task models.Task - if err := database.DB.Where("id = ?", taskID).First(&task).Error; err != nil { + res := database.DB.Where("id = ?", taskID).Limit(1).Find(&task) + if res.Error != nil || res.RowsAffected == 0 { return } @@ -121,8 +122,8 @@ func (s *TaskLogService) CleanTaskLogs(taskID string) { deleted = result.RowsAffected case "count": var boundaryLog models.TaskLog - err := database.DB.Where("task_id = ?", taskID).Order("id DESC").Offset(config.Keep - 1).Limit(1).First(&boundaryLog).Error - if err == nil { + res := database.DB.Where("task_id = ?", taskID).Order("id DESC").Offset(config.Keep - 1).Limit(1).Find(&boundaryLog) + if res.Error == nil && res.RowsAffected > 0 { result := database.DB.Where("task_id = ? AND id < ?", taskID, boundaryLog.ID).Delete(&models.TaskLog{}) deleted = result.RowsAffected } diff --git a/internal/services/tasks/task_service.go b/internal/services/tasks/task_service.go index 25de8ed..3f95ff4 100644 --- a/internal/services/tasks/task_service.go +++ b/internal/services/tasks/task_service.go @@ -15,7 +15,8 @@ func NewTaskService() *TaskService { func (ts *TaskService) GetTaskBySourceID(sourceID string) *models.Task { var task models.Task - if err := database.DB.Where("source_id = ?", sourceID).First(&task).Error; err != nil { + res := database.DB.Where("source_id = ?", sourceID).Limit(1).Find(&task) + if res.Error != nil || res.RowsAffected == 0 { return nil } return &task @@ -91,7 +92,8 @@ func (ts *TaskService) GetTasksWithPagination(page, pageSize int, name string, a func (ts *TaskService) GetTaskByID(id string) *models.Task { var task models.Task - if err := database.DB.Where("id = ?", id).First(&task).Error; err != nil { + res := database.DB.Where("id = ?", id).Limit(1).Find(&task) + if res.Error != nil || res.RowsAffected == 0 { return nil } return &task @@ -99,7 +101,8 @@ func (ts *TaskService) GetTaskByID(id string) *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, sourceID string) *models.Task { var task models.Task - if err := database.DB.Where("id = ?", id).First(&task).Error; err != nil { + res := database.DB.Where("id = ?", id).Limit(1).Find(&task) + if res.Error != nil || res.RowsAffected == 0 { return nil } task.Name = name diff --git a/internal/services/user_service.go b/internal/services/user_service.go index 2f4312d..b48581c 100644 --- a/internal/services/user_service.go +++ b/internal/services/user_service.go @@ -51,7 +51,8 @@ func (us *UserService) CreateUser(username, password, email, role string) *model func (us *UserService) GetUserByUsername(username string) *models.User { var user models.User - if err := database.DB.Where("username = ?", username).First(&user).Error; err != nil { + res := database.DB.Where("username = ?", username).Limit(1).Find(&user) + if res.Error != nil || res.RowsAffected == 0 { return nil } return &user @@ -59,8 +60,9 @@ func (us *UserService) GetUserByUsername(username string) *models.User { func (us *UserService) GetUserByID(id string) (*models.User, error) { var user models.User - if err := database.DB.Where("id = ?", id).First(&user).Error; err != nil { - return nil, err + res := database.DB.Where("id = ?", id).Limit(1).Find(&user) + if res.Error != nil || res.RowsAffected == 0 { + return nil, res.Error } return &user, nil } @@ -122,8 +124,9 @@ func (us *UserService) InvalidateUserTokens(userID string) error { func (us *UserService) UpdateAccount(userID string, newUsername string) error { var user models.User - if err := database.DB.Where("id = ?", userID).First(&user).Error; err != nil { - return err + res := database.DB.Where("id = ?", userID).Limit(1).Find(&user) + if res.Error != nil || res.RowsAffected == 0 { + return fmt.Errorf("未找到对应的用户") } updates := make(map[string]interface{})