Files

170 lines
4.9 KiB
Go
Raw Permalink Normal View History

2026-04-09 10:01:23 +08:00
package aiagent
import (
"code.tczkiot.com/wlw/ai-agent/cmd/testdata/seedlang"
"code.tczkiot.com/wlw/ai-agent/cmd/testdata/seeds"
"code.tczkiot.com/wlw/ai-agent/internal/models"
"code.tczkiot.com/wlw/ai-agent/internal/pkg/enums"
"code.tczkiot.com/wlw/ai-agent/internal/pkg/utils"
"code.tczkiot.com/wlw/ai-agent/internal/repositories"
2026-04-09 10:01:23 +08:00
"fmt"
"strings"
2026-04-09 10:01:23 +08:00
"time"
"github.com/mlogclub/simple/sqls"
)
type InitResult struct {
Created int
Updated int
}
// Init 初始化 AI Agent 测试数据
// 依赖于 AI Config 和 Knowledge Base 已初始化
func Init(lang seedlang.Language) (*InitResult, error) {
2026-04-09 10:01:23 +08:00
result := &InitResult{}
aiConfigID, err := getDefaultAIConfigID()
if err != nil {
return result, fmt.Errorf("get default ai config id failed: %w", err)
}
if aiConfigID == 0 {
return result, fmt.Errorf("no default ai config found, please init ai config first")
}
knowledgeIDs, err := getDefaultKnowledgeIDs()
if err != nil {
return result, fmt.Errorf("get default knowledge ids failed: %w", err)
}
defaultTeamIDs := getDefaultTeamIDs()
agentSeeds := seeds.AIAgentSeeds(lang)
seedItems := buildModels(lang, aiConfigID, knowledgeIDs, defaultTeamIDs)
for index, item := range seedItems {
2026-04-09 10:01:23 +08:00
itemCopy := item
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
existing := repositories.AIAgentRepository.Take(ctx.Tx, "name = ?", itemCopy.Name)
if existing == nil {
for _, legacyName := range agentSeeds[index].LegacyNames {
legacyName = strings.TrimSpace(legacyName)
if legacyName == "" {
continue
}
existing = repositories.AIAgentRepository.Take(ctx.Tx, "name = ?", legacyName)
if existing != nil {
break
}
}
}
2026-04-09 10:01:23 +08:00
if existing != nil {
if err := ctx.Tx.Model(existing).Updates(seedUpdateColumns(itemCopy)).Error; err != nil {
2026-04-09 10:01:23 +08:00
return err
}
result.Updated++
} else {
if err := ctx.Tx.Create(&itemCopy).Error; err != nil {
return err
}
result.Created++
}
return nil
}); err != nil {
return nil, fmt.Errorf("upsert ai agent failed: %w", err)
}
}
return result, nil
}
func buildModels(lang seedlang.Language, aiConfigID int64, knowledgeIDs []int64, defaultTeamIDs string) []models.AIAgent {
2026-04-09 10:01:23 +08:00
now := time.Now()
seedItems := seeds.AIAgentSeeds(lang)
items := make([]models.AIAgent, 0, len(seedItems))
for _, seed := range seedItems {
items = append(items, models.AIAgent{
Name: seed.Name,
Description: seed.Description,
Status: enums.StatusOk,
AIConfigID: aiConfigID,
ServiceMode: seed.ServiceMode,
SystemPrompt: seed.SystemPrompt,
WelcomeMessage: seed.WelcomeMessage,
ReplyTimeoutSeconds: seed.ReplyTimeoutSeconds,
2026-04-09 10:01:23 +08:00
TeamIDs: defaultTeamIDs,
HandoffMode: seed.HandoffMode,
FallbackMode: seed.FallbackMode,
FallbackMessage: seed.FallbackMessage,
2026-04-09 10:01:23 +08:00
KnowledgeIDs: utils.JoinInt64s(knowledgeIDs),
SortNo: seed.SortNo,
2026-04-09 10:01:23 +08:00
AuditFields: models.AuditFields{
CreatedAt: now,
CreateUserID: 0,
CreateUserName: "System",
UpdatedAt: now,
UpdateUserID: 0,
UpdateUserName: "System",
},
})
2026-04-09 10:01:23 +08:00
}
return items
2026-04-09 10:01:23 +08:00
}
func seedUpdateColumns(item models.AIAgent) map[string]any {
return map[string]any{
"name": item.Name,
"description": item.Description,
"status": item.Status,
"ai_config_id": item.AIConfigID,
"service_mode": item.ServiceMode,
"system_prompt": item.SystemPrompt,
"welcome_message": item.WelcomeMessage,
"reply_timeout_seconds": item.ReplyTimeoutSeconds,
"team_ids": item.TeamIDs,
"handoff_mode": item.HandoffMode,
"fallback_mode": item.FallbackMode,
"fallback_message": item.FallbackMessage,
"knowledge_ids": item.KnowledgeIDs,
"sort_no": item.SortNo,
"updated_at": item.UpdatedAt,
"update_user_id": item.UpdateUserID,
"update_user_name": item.UpdateUserName,
}
}
2026-04-09 10:01:23 +08:00
func getDefaultAIConfigID() (int64, error) {
aiConfig := repositories.AIConfigRepository.Take(
sqls.DB(),
"model_type = ? AND status = ?",
string(enums.AIModelTypeLLM),
enums.StatusOk,
)
if aiConfig == nil {
return 0, nil
}
return aiConfig.ID, nil
}
func getDefaultKnowledgeIDs() ([]int64, error) {
knowledges := repositories.KnowledgeBaseRepository.Find(
sqls.DB(),
sqls.NewCnd().Where("status = ?", enums.StatusOk),
)
ids := make([]int64, 0, len(knowledges))
for _, knowledge := range knowledges {
ids = append(ids, knowledge.ID)
}
return ids, nil
}
func getDefaultTeamIDs() string {
teams := repositories.AgentTeamRepository.Find(
sqls.DB(),
sqls.NewCnd().Where("status = ?", enums.StatusOk),
)
teamIDs := make([]int64, 0, len(teams))
for _, team := range teams {
teamIDs = append(teamIDs, team.ID)
}
return utils.JoinInt64s(teamIDs)
}