From 81056a61a187c2776cb323a2bf0636bb4c6fc9b6 Mon Sep 17 00:00:00 2001 From: mlogclub Date: Fri, 17 Apr 2026 17:27:33 +0800 Subject: [PATCH] feat(skill): optimize loadCandidateSkills to use batch retrieval for skill definitions --- internal/ai/skills/candidate_loader.go | 15 +++------------ .../repositories/skill_definition_repository.go | 15 +++++++++++++++ 2 files changed, 18 insertions(+), 12 deletions(-) diff --git a/internal/ai/skills/candidate_loader.go b/internal/ai/skills/candidate_loader.go index 4cbadde..a4ad80a 100644 --- a/internal/ai/skills/candidate_loader.go +++ b/internal/ai/skills/candidate_loader.go @@ -25,21 +25,12 @@ func (l *candidateLoader) loadCandidateSkills(aiAgent *models.AIAgent) []models. return nil } skillIDs := utils.SplitInt64s(aiAgent.SkillIDs) - if len(skillIDs) == 0 { - return nil - } + skills := repositories.SkillDefinitionRepository.GetByIDs(sqls.DB(), skillIDs) ret := make([]models.SkillDefinition, 0, len(skillIDs)) for _, id := range skillIDs { - // TODO 这里批量查询一下,批量查询返回数据的顺序需要保证和skillIDs一致 - skill := l.getSkillByID(id) - if skill == nil || skill.Status != enums.StatusOk { - continue + if skill, ok := skills[id]; ok && skill.Status == enums.StatusOk { + ret = append(ret, skill) } - ret = append(ret, *skill) } return ret } - -func (l *candidateLoader) getSkillByID(id int64) *models.SkillDefinition { - return repositories.SkillDefinitionRepository.Get(sqls.DB(), id) -} diff --git a/internal/repositories/skill_definition_repository.go b/internal/repositories/skill_definition_repository.go index dc94af4..0dc01c5 100644 --- a/internal/repositories/skill_definition_repository.go +++ b/internal/repositories/skill_definition_repository.go @@ -103,3 +103,18 @@ func (r *skillDefinitionRepository) Delete(db *gorm.DB, id int64) { func (r *skillDefinitionRepository) GetByCode(db *gorm.DB, code string) *models.SkillDefinition { return r.FindOne(db, sqls.NewCnd().Where("code = ?", code)) } + +func (r *skillDefinitionRepository) GetByIDs(db *gorm.DB, ids []int64) map[int64]models.SkillDefinition { + if len(ids) == 0 { + return nil + } + list := r.Find(db, sqls.NewCnd().Where("id IN (?)", ids)) + if len(list) == 0 { + return nil + } + result := make(map[int64]models.SkillDefinition, len(list)) + for _, item := range list { + result[item.ID] = item + } + return result +}