refactor: remove knowledge base references from AI agent model and related services

- Removed KnowledgeIDs from CreateAIAgentRequest and AIAgentResponse.
- Updated buildAIAgentResponseWithLocale to eliminate knowledge base name retrieval.
- Refactored AI agent repository and service to remove knowledge base handling.
- Adjusted dashboard service to no longer track AI agents without knowledge bases.
- Modified knowledge base service to check for workflow references instead of AI agent references.
- Updated frontend components to remove knowledge base selection and display.
- Enhanced workflow validation to ensure knowledge retrieve nodes have associated knowledge bases.
This commit is contained in:
mlogclub
2026-06-30 19:47:29 +08:00
parent fa0010e8be
commit 805ef87278
23 changed files with 575 additions and 243 deletions
-32
View File
@@ -110,7 +110,6 @@ func (s *aIAgentService) UpdateAIAgent(req request.UpdateAIAgentRequest, operato
"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,
"update_user_id": operator.UserID,
@@ -177,13 +176,6 @@ func (s *aIAgentService) buildAIAgentModel(id int64, req request.CreateAIAgentRe
return nil, errorsx.InvalidParamI18n("error.e0144")
}
knowledgeIDs, err := s.normalizeKnowledgeIDs(req.KnowledgeIDs)
if err != nil {
return nil, err
}
if len(knowledgeIDs) == 0 {
return nil, errorsx.InvalidParamI18n("error.e0320")
}
skillIDs, err := s.normalizeSkillIDs(req.SkillIDs)
if err != nil {
return nil, err
@@ -212,7 +204,6 @@ func (s *aIAgentService) buildAIAgentModel(id int64, req request.CreateAIAgentRe
HandoffMode: req.HandoffMode,
FallbackMode: req.FallbackMode,
FallbackMessage: strings.TrimSpace(req.FallbackMessage),
KnowledgeIDs: utils.JoinInt64s(knowledgeIDs),
SkillIDs: utils.JoinInt64s(skillIDs),
AllowedMCPTools: directToolsJSON,
WorkflowVersionID: 0,
@@ -243,29 +234,6 @@ func (s *aIAgentService) normalizeTeamIDs(input []int64) ([]int64, error) {
return ret, nil
}
func (s *aIAgentService) normalizeKnowledgeIDs(input []int64) ([]int64, error) {
ret := make([]int64, 0, len(input))
seen := make(map[int64]struct{})
for _, id := range input {
if id <= 0 {
continue
}
if _, exists := seen[id]; exists {
continue
}
kb := KnowledgeBaseService.Get(id)
if kb == nil || kb.Status == enums.StatusDeleted {
continue
}
// if kb.Status != enums.StatusOk {
// return nil, errorsx.InvalidParamI18n("error.e0285")
// }
seen[id] = struct{}{}
ret = append(ret, id)
}
return ret, nil
}
func (s *aIAgentService) normalizeSkillIDs(input []int64) ([]int64, error) {
ret := make([]int64, 0, len(input))
seen := make(map[int64]struct{})
@@ -22,7 +22,6 @@ func TestAIAgentServiceCreatesDefaultWorkflow(t *testing.T) {
setupAIAgentWorkflowTestDB(t)
operator := aiAgentWorkflowTestOperator()
aiConfigID := createAIAgentWorkflowTestConfig(t)
knowledgeID := createAIAgentWorkflowTestKnowledgeBase(t)
item, err := AIAgentService.CreateAIAgent(request.CreateAIAgentRequest{
Name: "workflow agent",
@@ -30,7 +29,6 @@ func TestAIAgentServiceCreatesDefaultWorkflow(t *testing.T) {
ServiceMode: enums.IMConversationServiceModeAIOnly,
HandoffMode: enums.AIAgentHandoffModeWaitPool,
FallbackMode: enums.AIAgentFallbackModeNoAnswer,
KnowledgeIDs: []int64{knowledgeID},
}, operator)
if err != nil {
t.Fatalf("CreateAIAgent() error = %v", err)
@@ -54,8 +52,8 @@ func TestAIAgentServiceCreatesDefaultWorkflow(t *testing.T) {
t.Fatalf("expected default draft definition")
}
validation := workflowvalidator.ValidateDefinition(stored, workflowregistry.DefaultRegistry())
if !validation.Valid {
t.Fatalf("expected default workflow to be valid, got %#v", validation.Errors)
if validation.Valid || !workflowValidationHasMessage(validation, "需要选择至少一个知识库") {
t.Fatalf("expected default workflow to require node knowledge bases, got %#v", validation.Errors)
}
if nodeTypeByID(stored, "understanding_1") != workflowregistry.NodeTypeConversationUnderstanding {
t.Fatalf("expected default workflow to include conversation understanding, got nodes: %#v", stored.Nodes)
@@ -116,14 +114,14 @@ func TestAIAgentServiceCreatesDefaultWorkflow(t *testing.T) {
})
}
func TestAIWorkflowServiceDefaultAgentWorkflowDefinitionIsValid(t *testing.T) {
func TestAIWorkflowServiceDefaultAgentWorkflowDefinitionRequiresKnowledgeRetrieveConfiguration(t *testing.T) {
definition := AIWorkflowService.DefaultAgentWorkflowDefinition()
if definition.SchemaVersion != dsl.SchemaVersion || nodeTypeByID(definition, "start_1") != workflowregistry.NodeTypeStart {
t.Fatalf("expected default workflow definition")
}
validation := workflowvalidator.ValidateDefinition(definition, workflowregistry.DefaultRegistry())
if !validation.Valid {
t.Fatalf("expected default workflow definition to be valid, got %#v", validation.Errors)
if validation.Valid || !workflowValidationHasMessage(validation, "需要选择至少一个知识库") {
t.Fatalf("expected default workflow definition to require node knowledge bases, got %#v", validation.Errors)
}
if nodeTypeByID(definition, "understanding_1") != workflowregistry.NodeTypeConversationUnderstanding {
t.Fatalf("expected default workflow to include conversation understanding, got nodes: %#v", definition.Nodes)
@@ -167,7 +165,6 @@ func TestAIWorkflowServicePublishAgentWorkflowBindsAgentVersion(t *testing.T) {
setupAIAgentWorkflowTestDB(t)
operator := aiAgentWorkflowTestOperator()
aiConfigID := createAIAgentWorkflowTestConfig(t)
knowledgeID := createAIAgentWorkflowTestKnowledgeBase(t)
agent, err := AIAgentService.CreateAIAgent(request.CreateAIAgentRequest{
Name: "workflow agent without version",
@@ -175,7 +172,6 @@ func TestAIWorkflowServicePublishAgentWorkflowBindsAgentVersion(t *testing.T) {
ServiceMode: enums.IMConversationServiceModeAIOnly,
HandoffMode: enums.AIAgentHandoffModeWaitPool,
FallbackMode: enums.AIAgentFallbackModeNoAnswer,
KnowledgeIDs: []int64{knowledgeID},
}, operator)
if err != nil {
t.Fatalf("CreateAIAgent() error = %v", err)
@@ -287,6 +283,15 @@ func workflowHasNodeType(def dsl.Definition, nodeType string) bool {
return false
}
func workflowValidationHasMessage(result workflowvalidator.Result, message string) bool {
for _, item := range result.Errors {
if strings.Contains(item.Message, message) {
return true
}
}
return false
}
func nodeTypeByID(def dsl.Definition, nodeID string) string {
for _, node := range def.Nodes {
if node.ID == nodeID {
+1 -1
View File
@@ -472,7 +472,7 @@ func defaultAgentWorkflowDefinition() dsl.Definition {
"followUpQuestions": dsl.RefValue("draft_ticket_1", "followUpQuestions"),
}, map[string]any{"staticReply": "为了创建工单,还需要补充以下信息:\n{{followUpQuestions}}"}),
workflowNode("send_ticket_followup_1", workflowregistry.NodeTypeSendReply, "发送工单追问", 4780, 1033.5, workflowInputs("replyText", "ticket_followup_reply_1", "replyText"), nil),
workflowNode("retrieve_1", workflowregistry.NodeTypeKnowledgeRetrieve, "知识检索", 2480, 753, workflowInputs("query", "start_1", "userMessage"), nil),
workflowNode("retrieve_1", workflowregistry.NodeTypeKnowledgeRetrieve, "知识检索", 2480, 753, workflowInputs("query", "start_1", "userMessage"), map[string]any{"knowledgeBaseIds": []int64{}}),
workflowNode("answerability_1", workflowregistry.NodeTypeAnswerabilityGate, "可回答判断", 2940, 753, map[string]dsl.Value{
"userMessage": dsl.RefValue("start_1", "userMessage"),
"knowledgeItems": dsl.RefValue("retrieve_1", "items"),
-17
View File
@@ -274,23 +274,6 @@ func (s *dashboardService) buildAlerts(now time.Time, db *gorm.DB, aiAgents []mo
})
}
var aiAgentWithoutKnowledgeCount int64
for _, item := range aiAgents {
if strings.TrimSpace(item.KnowledgeIDs) == "" {
aiAgentWithoutKnowledgeCount++
}
}
if aiAgentWithoutKnowledgeCount > 0 {
alerts = append(alerts, response.DashboardAlertResponse{
ID: "ai-no-knowledge",
Level: "info",
Title: dashboardText(locale, "alert.aiNoKnowledge.title"),
Description: dashboardText(locale, "alert.aiNoKnowledge.description"),
Count: aiAgentWithoutKnowledgeCount,
Link: "/dashboard/ai-agents",
})
}
sort.Slice(alerts, func(i, j int) bool {
if alerts[i].Count == alerts[j].Count {
return alerts[i].ID < alerts[j].ID
+90 -5
View File
@@ -2,9 +2,14 @@ package services
import (
"context"
"encoding/json"
"fmt"
"strings"
"time"
"agent-desk/internal/ai/rag"
"agent-desk/internal/ai/workflow/dsl"
workflowregistry "agent-desk/internal/ai/workflow/registry"
"agent-desk/internal/models"
"agent-desk/internal/pkg/dto"
"agent-desk/internal/pkg/dto/request"
@@ -128,12 +133,12 @@ func (s *knowledgeBaseService) DeleteKnowledgeBase(id int64) error {
return errorsx.InvalidParamI18n("error.e0283")
}
referencingAgents := repositories.AIAgentRepository.FindByKnowledgeBaseID(sqls.DB(), id)
if len(referencingAgents) > 0 {
if len(referencingAgents) == 1 {
return errorsx.ForbiddenI18n("error.knowledgeBase.referencedByAgent", referencingAgents[0].Name)
referencingWorkflows := s.findWorkflowReferencesByKnowledgeBaseID(id)
if len(referencingWorkflows) > 0 {
if len(referencingWorkflows) == 1 {
return errorsx.Forbidden(fmt.Sprintf("知识库正在被流程「%s」使用,请先从知识检索节点中移除", referencingWorkflows[0]))
}
return errorsx.ForbiddenI18n("error.knowledgeBase.referencedByAgents", len(referencingAgents))
return errorsx.Forbidden(fmt.Sprintf("知识库正在被 %d 个流程使用,请先从知识检索节点中移除", len(referencingWorkflows)))
}
if err := sqls.WithTransaction(func(ctx *sqls.TxContext) error {
@@ -151,6 +156,86 @@ func (s *knowledgeBaseService) DeleteKnowledgeBase(id int64) error {
return rag.Index.RemoveKnowledgeBaseIndex(context.Background(), id)
}
func (s *knowledgeBaseService) findWorkflowReferencesByKnowledgeBaseID(id int64) []string {
names := make(map[string]struct{})
workflows := repositories.AIWorkflowRepository.Find(sqls.DB(), sqls.NewCnd().Eq("status", enums.StatusOk))
workflowNames := make(map[int64]string, len(workflows))
for _, workflow := range workflows {
name := strings.TrimSpace(workflow.Name)
if name == "" {
name = fmt.Sprintf("ID %d", workflow.ID)
}
workflowNames[workflow.ID] = name
if workflowDefinitionUsesKnowledgeBase(workflow.DraftDefinition, id) {
names[name] = struct{}{}
}
}
versions := repositories.AIWorkflowVersionRepository.Find(sqls.DB(), sqls.NewCnd())
for _, version := range versions {
if !workflowDefinitionUsesKnowledgeBase(version.Definition, id) {
continue
}
name := workflowNames[version.WorkflowID]
if strings.TrimSpace(name) == "" {
name = fmt.Sprintf("ID %d", version.WorkflowID)
}
names[name] = struct{}{}
}
ret := make([]string, 0, len(names))
for name := range names {
ret = append(ret, name)
}
return ret
}
func workflowDefinitionUsesKnowledgeBase(definition string, id int64) bool {
definition = strings.TrimSpace(definition)
if definition == "" {
return false
}
var def dsl.Definition
if err := json.Unmarshal([]byte(definition), &def); err != nil {
return false
}
for _, node := range def.Nodes {
if strings.TrimSpace(node.Type) != workflowregistry.NodeTypeKnowledgeRetrieve {
continue
}
for _, knowledgeBaseID := range knowledgeBaseIDsFromWorkflowNodeConfig(node.Data.Config) {
if knowledgeBaseID == id {
return true
}
}
}
return false
}
func knowledgeBaseIDsFromWorkflowNodeConfig(raw json.RawMessage) []int64 {
if len(raw) == 0 {
return nil
}
var cfg map[string]any
if err := json.Unmarshal(raw, &cfg); err != nil {
return nil
}
items, ok := cfg["knowledgeBaseIds"].([]any)
if !ok {
return nil
}
ret := make([]int64, 0, len(items))
for _, item := range items {
switch value := item.(type) {
case float64:
ret = append(ret, int64(value))
case int64:
ret = append(ret, value)
case int:
ret = append(ret, int64(value))
}
}
return ret
}
func (s *knowledgeBaseService) UpdateSort(ids []int64) error {
return sqls.WithTransaction(func(ctx *sqls.TxContext) error {
for i, id := range ids {
@@ -1,10 +1,12 @@
package services
import (
"fmt"
"encoding/json"
"strings"
"testing"
"agent-desk/internal/ai/workflow/dsl"
workflowregistry "agent-desk/internal/ai/workflow/registry"
"agent-desk/internal/models"
"agent-desk/internal/pkg/dto/request"
"agent-desk/internal/pkg/enums"
@@ -25,34 +27,48 @@ func TestBuildKnowledgeBaseModelUsesLowerDefaultScoreThreshold(t *testing.T) {
}
}
func TestDeleteKnowledgeBaseRejectsAIAgentReference(t *testing.T) {
func TestDeleteKnowledgeBaseRejectsWorkflowDraftReference(t *testing.T) {
setupKnowledgeBaseServiceTestDB(t)
kb := createKnowledgeBaseServiceTestBase(t, "Referenced KB")
otherKB := createKnowledgeBaseServiceTestBase(t, "Other KB")
if err := repositories.AIAgentRepository.Create(sqls.DB(), &models.AIAgent{
Name: "Support Agent",
Status: enums.StatusOk,
KnowledgeIDs: "12",
}); err != nil {
t.Fatalf("create unrelated ai agent: %v", err)
}
if err := repositories.AIAgentRepository.Create(sqls.DB(), &models.AIAgent{
Name: "Knowledge Agent",
Status: enums.StatusOk,
KnowledgeIDs: fmt.Sprintf("12,%d,%d", kb.ID, otherKB.ID),
}); err != nil {
t.Fatalf("create ai agent: %v", err)
}
createKnowledgeBaseServiceTestWorkflow(t, "Support Workflow", knowledgeBaseServiceTestWorkflowDefinition([]int64{12, otherKB.ID}))
createKnowledgeBaseServiceTestWorkflow(t, "Knowledge Workflow", knowledgeBaseServiceTestWorkflowDefinition([]int64{12, kb.ID, otherKB.ID}))
err := KnowledgeBaseService.DeleteKnowledgeBase(kb.ID)
if err == nil {
t.Fatal("DeleteKnowledgeBase() error is nil, want referenced knowledge base error")
t.Fatal("DeleteKnowledgeBase() error is nil, want referenced workflow error")
}
if got := err.Error(); !strings.Contains(got, "Knowledge Agent") {
t.Fatalf("DeleteKnowledgeBase() error = %q, want agent name", got)
if got := err.Error(); !strings.Contains(got, "Knowledge Workflow") {
t.Fatalf("DeleteKnowledgeBase() error = %q, want workflow name", got)
}
if repositories.KnowledgeBaseRepository.Get(sqls.DB(), kb.ID) == nil {
t.Fatal("knowledge base was deleted despite ai agent reference")
t.Fatal("knowledge base was deleted despite workflow reference")
}
}
func TestDeleteKnowledgeBaseRejectsWorkflowVersionReference(t *testing.T) {
setupKnowledgeBaseServiceTestDB(t)
kb := createKnowledgeBaseServiceTestBase(t, "Version KB")
workflow := createKnowledgeBaseServiceTestWorkflow(t, "Published Workflow", knowledgeBaseServiceTestWorkflowDefinition([]int64{999}))
raw, err := json.Marshal(knowledgeBaseServiceTestWorkflowDefinition([]int64{kb.ID}))
if err != nil {
t.Fatalf("marshal workflow version definition: %v", err)
}
if err := repositories.AIWorkflowVersionRepository.Create(sqls.DB(), &models.AIWorkflowVersion{
WorkflowID: workflow.ID,
Version: 1,
Status: enums.StatusOk,
Definition: string(raw),
}); err != nil {
t.Fatalf("create workflow version: %v", err)
}
err = KnowledgeBaseService.DeleteKnowledgeBase(kb.ID)
if err == nil {
t.Fatal("DeleteKnowledgeBase() error is nil, want referenced workflow version error")
}
if got := err.Error(); !strings.Contains(got, "Published Workflow") {
t.Fatalf("DeleteKnowledgeBase() error = %q, want workflow name", got)
}
}
@@ -103,12 +119,67 @@ func setupKnowledgeBaseServiceTestDB(t *testing.T) {
if err != nil {
t.Fatalf("open sqlite db: %v", err)
}
if err := db.AutoMigrate(&models.KnowledgeBase{}, &models.KnowledgeDocument{}, &models.KnowledgeFAQ{}, &models.KnowledgeChunk{}, &models.AIAgent{}); err != nil {
if err := db.AutoMigrate(&models.KnowledgeBase{}, &models.KnowledgeDocument{}, &models.KnowledgeFAQ{}, &models.KnowledgeChunk{}, &models.AIAgent{}, &models.AIWorkflow{}, &models.AIWorkflowVersion{}); err != nil {
t.Fatalf("auto migrate: %v", err)
}
sqls.SetDB(db)
}
func createKnowledgeBaseServiceTestWorkflow(t *testing.T, name string, definition dsl.Definition) *models.AIWorkflow {
t.Helper()
raw, err := json.Marshal(definition)
if err != nil {
t.Fatalf("marshal workflow definition: %v", err)
}
item := &models.AIWorkflow{
Name: name,
Status: enums.StatusOk,
DraftDefinition: string(raw),
}
if err := repositories.AIWorkflowRepository.Create(sqls.DB(), item); err != nil {
t.Fatalf("create workflow: %v", err)
}
return item
}
func knowledgeBaseServiceTestWorkflowDefinition(knowledgeBaseIDs []int64) dsl.Definition {
return dsl.Definition{
SchemaVersion: dsl.SchemaVersion,
Nodes: []dsl.Node{
{
ID: "start_1",
Type: workflowregistry.NodeTypeStart,
},
{
ID: "retrieve_1",
Type: workflowregistry.NodeTypeKnowledgeRetrieve,
Data: dsl.NodeData{
Config: mustKnowledgeBaseServiceTestJSON(map[string]any{"knowledgeBaseIds": knowledgeBaseIDs}),
InputsValues: map[string]dsl.Value{
"query": dsl.RefValue("start_1", "userMessage"),
},
},
},
{
ID: "end_1",
Type: workflowregistry.NodeTypeEnd,
},
},
Edges: []dsl.Edge{
{SourceNodeID: "start_1", TargetNodeID: "retrieve_1"},
{SourceNodeID: "retrieve_1", TargetNodeID: "end_1"},
},
}
}
func mustKnowledgeBaseServiceTestJSON(value any) json.RawMessage {
raw, err := json.Marshal(value)
if err != nil {
panic(err)
}
return raw
}
func createKnowledgeBaseServiceTestBase(t *testing.T, name string) *models.KnowledgeBase {
t.Helper()
item := &models.KnowledgeBase{