refactor: enhance AIAgent seed structure and update initialization logic for improved agent handling
This commit is contained in:
Vendored
+44
-24
@@ -3,12 +3,12 @@ package aiagent
|
||||
import (
|
||||
"agent-desk/cmd/testdata/seedlang"
|
||||
"agent-desk/cmd/testdata/seeds"
|
||||
"agent-desk/cmd/testdata/skill"
|
||||
"agent-desk/internal/models"
|
||||
"agent-desk/internal/pkg/enums"
|
||||
"agent-desk/internal/pkg/utils"
|
||||
"agent-desk/internal/repositories"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mlogclub/simple/sqls"
|
||||
@@ -38,24 +38,30 @@ func Init(lang seedlang.Language) (*InitResult, error) {
|
||||
}
|
||||
|
||||
defaultTeamIDs := getDefaultTeamIDs()
|
||||
defaultSkillIDs, err := getDefaultSkillIDs()
|
||||
if err != nil {
|
||||
return result, fmt.Errorf("get default skill ids failed: %w", err)
|
||||
}
|
||||
|
||||
seedItems := buildModels(lang, aiConfigID, knowledgeIDs, defaultTeamIDs, defaultSkillIDs)
|
||||
for _, item := range seedItems {
|
||||
agentSeeds := seeds.AIAgentSeeds(lang)
|
||||
seedItems := buildModels(lang, aiConfigID, knowledgeIDs, defaultTeamIDs)
|
||||
for index, item := range seedItems {
|
||||
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
|
||||
}
|
||||
}
|
||||
}
|
||||
if existing != nil {
|
||||
// 更新
|
||||
if err := ctx.Tx.Model(existing).Updates(&itemCopy).Error; err != nil {
|
||||
if err := ctx.Tx.Model(existing).Updates(seedUpdateColumns(itemCopy)).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
result.Updated++
|
||||
} else {
|
||||
// 创建
|
||||
if err := ctx.Tx.Create(&itemCopy).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -70,7 +76,7 @@ func Init(lang seedlang.Language) (*InitResult, error) {
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func buildModels(lang seedlang.Language, aiConfigID int64, knowledgeIDs []int64, defaultTeamIDs string, defaultSkillIDs string) []models.AIAgent {
|
||||
func buildModels(lang seedlang.Language, aiConfigID int64, knowledgeIDs []int64, defaultTeamIDs string) []models.AIAgent {
|
||||
now := time.Now()
|
||||
seedItems := seeds.AIAgentSeeds(lang)
|
||||
items := make([]models.AIAgent, 0, len(seedItems))
|
||||
@@ -89,7 +95,8 @@ func buildModels(lang seedlang.Language, aiConfigID int64, knowledgeIDs []int64,
|
||||
FallbackMode: seed.FallbackMode,
|
||||
FallbackMessage: seed.FallbackMessage,
|
||||
KnowledgeIDs: utils.JoinInt64s(knowledgeIDs),
|
||||
SkillIDs: defaultSkillIDs,
|
||||
SkillIDs: "",
|
||||
AllowedMCPTools: "",
|
||||
SortNo: seed.SortNo,
|
||||
AuditFields: models.AuditFields{
|
||||
CreatedAt: now,
|
||||
@@ -104,6 +111,30 @@ func buildModels(lang seedlang.Language, aiConfigID int64, knowledgeIDs []int64,
|
||||
return items
|
||||
}
|
||||
|
||||
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,
|
||||
"skill_ids": item.SkillIDs,
|
||||
"allowed_mcp_tools": item.AllowedMCPTools,
|
||||
"sort_no": item.SortNo,
|
||||
"updated_at": item.UpdatedAt,
|
||||
"update_user_id": item.UpdateUserID,
|
||||
"update_user_name": item.UpdateUserName,
|
||||
}
|
||||
}
|
||||
|
||||
func getDefaultAIConfigID() (int64, error) {
|
||||
aiConfig := repositories.AIConfigRepository.Take(
|
||||
sqls.DB(),
|
||||
@@ -140,14 +171,3 @@ func getDefaultTeamIDs() string {
|
||||
}
|
||||
return utils.JoinInt64s(teamIDs)
|
||||
}
|
||||
|
||||
func getDefaultSkillIDs() (string, error) {
|
||||
skillItem := repositories.SkillDefinitionRepository.FindOne(
|
||||
sqls.DB(),
|
||||
sqls.NewCnd().Where("status = ?", enums.StatusOk).Desc("id"),
|
||||
)
|
||||
if skillItem == nil {
|
||||
return "", fmt.Errorf("default test skill not found")
|
||||
}
|
||||
return utils.JoinInt64s([]int64{skillItem.ID}), nil
|
||||
}
|
||||
|
||||
Vendored
+56
@@ -4,6 +4,7 @@ import (
|
||||
"agent-desk/cmd/testdata/seedlang"
|
||||
"agent-desk/cmd/testdata/seeds"
|
||||
"regexp"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
@@ -12,6 +13,7 @@ var hanTextPattern = regexp.MustCompile(`\p{Han}`)
|
||||
func TestEnglishAIAgentSeedDoesNotContainChineseText(t *testing.T) {
|
||||
for _, item := range seeds.AIAgentSeeds(seedlang.English) {
|
||||
values := []string{item.Name, item.Description, item.SystemPrompt, item.WelcomeMessage, item.FallbackMessage}
|
||||
values = append(values, item.LegacyNames...)
|
||||
for _, value := range values {
|
||||
if hanTextPattern.MatchString(value) {
|
||||
t.Fatalf("english AI agent seed contains Chinese text: %q", value)
|
||||
@@ -19,3 +21,57 @@ func TestEnglishAIAgentSeedDoesNotContainChineseText(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestChineseAIAgentSeedUsesPresalesConfiguration(t *testing.T) {
|
||||
items := seeds.AIAgentSeeds(seedlang.Chinese)
|
||||
if len(items) != 1 {
|
||||
t.Fatalf("expected one Chinese AI agent seed, got %d", len(items))
|
||||
}
|
||||
|
||||
item := items[0]
|
||||
if item.Name != "贝壳AI售前客服" {
|
||||
t.Fatalf("unexpected Chinese AI agent name: %q", item.Name)
|
||||
}
|
||||
for _, expected := range []string{
|
||||
"你是贝壳AI(AgentDesk)的售前客服",
|
||||
"# 售前对话策略",
|
||||
"# 事实与工具边界",
|
||||
"不得编造价格、折扣、合同条款、商业 SLA",
|
||||
} {
|
||||
if !strings.Contains(item.SystemPrompt, expected) {
|
||||
t.Fatalf("Chinese AI agent system prompt missing %q", expected)
|
||||
}
|
||||
}
|
||||
if !strings.Contains(item.WelcomeMessage, "贝壳AI售前客服") {
|
||||
t.Fatalf("unexpected welcome message: %q", item.WelcomeMessage)
|
||||
}
|
||||
if !strings.Contains(item.FallbackMessage, "方案评估") {
|
||||
t.Fatalf("unexpected fallback message: %q", item.FallbackMessage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildModelsLeavesSkillsAndMCPToolsUnbound(t *testing.T) {
|
||||
items := buildModels(seedlang.Chinese, 7, []int64{11}, "13")
|
||||
if len(items) != 1 {
|
||||
t.Fatalf("expected one AI agent model, got %d", len(items))
|
||||
}
|
||||
|
||||
item := items[0]
|
||||
if item.AIConfigID != 7 || item.KnowledgeIDs != "11" || item.TeamIDs != "13" {
|
||||
t.Fatalf("unexpected AI agent bindings: %+v", item)
|
||||
}
|
||||
if item.SkillIDs != "" {
|
||||
t.Fatalf("expected no Skill binding, got %q", item.SkillIDs)
|
||||
}
|
||||
if item.AllowedMCPTools != "" {
|
||||
t.Fatalf("expected no MCP Tool binding, got %q", item.AllowedMCPTools)
|
||||
}
|
||||
|
||||
columns := seedUpdateColumns(item)
|
||||
if value, ok := columns["skill_ids"]; !ok || value != "" {
|
||||
t.Fatalf("seed update must clear Skill bindings, got %#v", value)
|
||||
}
|
||||
if value, ok := columns["allowed_mcp_tools"]; !ok || value != "" {
|
||||
t.Fatalf("seed update must clear MCP Tool bindings, got %#v", value)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user