From 805ef87278247e92b909ed4d563062311a5dc797 Mon Sep 17 00:00:00 2001 From: mlogclub Date: Tue, 30 Jun 2026 19:47:29 +0800 Subject: [PATCH] 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. --- .../runtime/retrievers/knowledge_retriever.go | 19 ++- internal/ai/runtime/workflow/executor.go | 36 +++- internal/ai/runtime/workflow/executor_test.go | 35 +++- internal/ai/workflow/registry/registry.go | 9 +- .../ai/workflow/registry/registry_test.go | 7 +- internal/ai/workflow/validator/validator.go | 53 ++++++ .../ai/workflow/validator/validator_test.go | 51 ++++++ .../handlers/dashboard/ai_agent_handler.go | 7 - internal/pkg/dto/request/ai_request.go | 1 - internal/pkg/dto/response/ai_response.go | 2 - internal/repositories/ai_agent_repository.go | 17 -- internal/services/ai_agent_service.go | 32 ---- .../ai_agent_workflow_service_test.go | 23 ++- internal/services/ai_workflow_service.go | 2 +- internal/services/dashboard_service.go | 17 -- internal/services/knowledge_base_service.go | 95 ++++++++++- .../services/knowledge_base_service_test.go | 113 ++++++++++--- .../_components/config-workbench.tsx | 81 --------- web/app/dashboard/ai-agents/page.tsx | 26 --- .../_components/node-config-panel.tsx | 159 +++++++++++++++++- .../_components/workflow-utils.test.mjs | 21 +++ .../_components/workflow-utils.ts | 9 + web/lib/api/admin.ts | 3 - 23 files changed, 575 insertions(+), 243 deletions(-) diff --git a/internal/ai/runtime/retrievers/knowledge_retriever.go b/internal/ai/runtime/retrievers/knowledge_retriever.go index 13704fb..661799b 100644 --- a/internal/ai/runtime/retrievers/knowledge_retriever.go +++ b/internal/ai/runtime/retrievers/knowledge_retriever.go @@ -8,7 +8,6 @@ import ( "agent-desk/internal/ai/runtime/traces" "agent-desk/internal/models" "agent-desk/internal/pkg/enums" - "agent-desk/internal/pkg/utils" "agent-desk/internal/repositories" "github.com/mlogclub/simple/sqls" @@ -20,7 +19,8 @@ const defaultRuntimeKnowledgeScoreThreshold = 0.3 const defaultRuntimeKnowledgeMaxContextItems = 5 type KnowledgeRetriever struct { - AIAgent models.AIAgent + AIAgent models.AIAgent + knowledgeBaseIDs []int64 } type KnowledgeRetrieveOptions struct { @@ -52,8 +52,11 @@ type KnowledgeRetrieveResult struct { Policies []KnowledgeBaseRetrievePolicy } -func NewKnowledgeRetriever(aiAgent models.AIAgent) *KnowledgeRetriever { - return &KnowledgeRetriever{AIAgent: aiAgent} +func NewKnowledgeRetriever(aiAgent models.AIAgent, knowledgeBaseIDs []int64) *KnowledgeRetriever { + return &KnowledgeRetriever{ + AIAgent: aiAgent, + knowledgeBaseIDs: append([]int64(nil), knowledgeBaseIDs...), + } } func DefaultKnowledgeRetrieveOptions() KnowledgeRetrieveOptions { @@ -63,8 +66,8 @@ func DefaultKnowledgeRetrieveOptions() KnowledgeRetrieveOptions { } } -func (r *KnowledgeRetriever) KnowledgeBaseIDs() []int64 { - return utils.SplitInt64s(r.AIAgent.KnowledgeIDs) +func (r *KnowledgeRetriever) ConfiguredKnowledgeBaseIDs() []int64 { + return append([]int64(nil), r.knowledgeBaseIDs...) } func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) { @@ -72,7 +75,7 @@ func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag. } func (r *KnowledgeRetriever) RetrieveByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) { - ids := r.KnowledgeBaseIDs() + ids := r.ConfiguredKnowledgeBaseIDs() return rag.Retrieve.RetrieveWithTrace(ctx, rag.RetrieveRequest{ Query: query, KnowledgeBaseIDs: ids, @@ -87,7 +90,7 @@ func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string) func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) (*KnowledgeRetrieveResult, error) { query = strings.TrimSpace(query) - knowledgeBaseIDs := r.KnowledgeBaseIDs() + knowledgeBaseIDs := r.ConfiguredKnowledgeBaseIDs() policies := r.resolvePolicies(knowledgeBaseIDs, opts) contextMaxTokens := opts.ContextMaxTokens if contextMaxTokens <= 0 { diff --git a/internal/ai/runtime/workflow/executor.go b/internal/ai/runtime/workflow/executor.go index 6b014e1..2d2e73b 100644 --- a/internal/ai/runtime/workflow/executor.go +++ b/internal/ai/runtime/workflow/executor.go @@ -20,7 +20,6 @@ import ( "agent-desk/internal/pkg/dto" "agent-desk/internal/pkg/dto/request" "agent-desk/internal/pkg/enums" - "agent-desk/internal/pkg/utils" "agent-desk/internal/services" ) @@ -270,7 +269,6 @@ func (e *Executor) executeNode(ctx context.Context, state *runState, node dsl.No "messageId": state.input.UserMessage.ID, "aiAgentId": state.input.AIAgent.ID, "userMessage": strings.TrimSpace(state.input.UserMessage.Content), - "knowledgeBaseIds": utils.SplitInt64s(state.input.AIAgent.KnowledgeIDs), "conversationState": state.input.Conversation.Status, }) case workflowregistry.NodeTypeConversationUnderstanding: @@ -732,7 +730,11 @@ func (e *Executor) executeHandoffToHuman(state *runState, node dsl.Node) error { func (e *Executor) executeKnowledgeRetrieve(ctx context.Context, state *runState, node dsl.Node) error { query := strings.TrimSpace(toString(state.resolveInput(node, "query"))) - retriever := retrievers.NewKnowledgeRetriever(state.input.AIAgent) + knowledgeBaseIDs := readInt64ArrayConfig(node.Data.Config, "knowledgeBaseIds") + if len(knowledgeBaseIDs) == 0 { + return fmt.Errorf("knowledge retrieve node requires knowledgeBaseIds") + } + retriever := retrievers.NewKnowledgeRetriever(state.input.AIAgent, knowledgeBaseIDs) result, err := retriever.RetrieveContext(ctx, query) if err != nil { return err @@ -784,7 +786,7 @@ func (e *Executor) executeLLMReply(ctx context.Context, state *runState, node ds if prompt := strings.TrimSpace(readStringConfig(node.Data.Config, "prompt")); prompt != "" { systemPrompt = strings.TrimSpace(systemPrompt + "\n\n" + prompt) } - if _, declaresKnowledge := node.Data.InputsValues["knowledgeItems"]; declaresKnowledge && len(utils.SplitInt64s(state.input.AIAgent.KnowledgeIDs)) > 0 && !hasItems(state.resolveInput(node, "knowledgeItems")) { + if _, declaresKnowledge := node.Data.InputsValues["knowledgeItems"]; declaresKnowledge && !hasItems(state.resolveInput(node, "knowledgeItems")) { state.setNodeVars(node.ID, map[string]any{"replyText": workflowKnowledgeFallbackReply(state.input.AIAgent)}) return nil } @@ -1088,6 +1090,32 @@ func readBoolConfig(raw json.RawMessage, key string) bool { return truthy(cfg[key]) } +func readInt64ArrayConfig(raw json.RawMessage, key string) []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[key].([]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 compareString(left any, right any) int { return strings.Compare(toString(left), toString(right)) } diff --git a/internal/ai/runtime/workflow/executor_test.go b/internal/ai/runtime/workflow/executor_test.go index f6adc8f..36cd80d 100644 --- a/internal/ai/runtime/workflow/executor_test.go +++ b/internal/ai/runtime/workflow/executor_test.go @@ -298,7 +298,6 @@ func TestExecutorPolicyFirstWorkflowRoutesGreetingToDirectReply(t *testing.T) { Content: "

你好。

", }, AIAgent: models.AIAgent{ - KnowledgeIDs: "1", FallbackMessage: "我暂时没有找到足够准确的信息。", }, }) @@ -330,7 +329,6 @@ func TestExecutorPolicyFirstWorkflowRoutesBusinessQuestionToKnowledge(t *testing Content: "你们价格是多少?", }, AIAgent: models.AIAgent{ - KnowledgeIDs: "1", FallbackMessage: "我暂时没有找到足够准确的信息。", }, }) @@ -340,6 +338,24 @@ func TestExecutorPolicyFirstWorkflowRoutesBusinessQuestionToKnowledge(t *testing assertPath(t, result.NodePath, []string{"start_1", "understanding_1", "policy_1", "policy_route_1", "retrieve_end"}) } +func TestExecutorKnowledgeRetrieveRequiresNodeKnowledgeBases(t *testing.T) { + _, err := NewExecutor().Execute(context.Background(), Input{ + Definition: knowledgeRetrieveWorkflowDefinition(nil), + UserMessage: models.Message{ + Content: "产品价格", + }, + AIAgent: models.AIAgent{ + KnowledgeIDs: "1,2", + }, + }) + if err == nil { + t.Fatalf("expected knowledge retrieve without node knowledge bases to fail") + } + if !strings.Contains(err.Error(), "knowledge retrieve node requires knowledgeBaseIds") { + t.Fatalf("unexpected error: %v", err) + } +} + func TestExecutorLLMReplyUsesAgentFallbackWhenDeclaredKnowledgeIsEmpty(t *testing.T) { result, err := NewExecutor().Execute(context.Background(), Input{ Definition: emptyKnowledgeReplyDefinition(), @@ -347,7 +363,6 @@ func TestExecutorLLMReplyUsesAgentFallbackWhenDeclaredKnowledgeIsEmpty(t *testin Content: "产品功能", }, AIAgent: models.AIAgent{ - KnowledgeIDs: "1", FallbackMode: enums.AIAgentFallbackModeNoAnswer, FallbackMessage: "我暂时没有找到足够准确的信息。你可以补充更具体的问题,我再继续帮你查。", SystemPrompt: "不要编造事实。", @@ -564,6 +579,20 @@ func conditionalReplyDefinition() dsl.Definition { ) } +func knowledgeRetrieveWorkflowDefinition(config any) dsl.Definition { + return wfTestDefinition( + []dsl.Node{ + wfTestNode("start_1", workflowregistry.NodeTypeStart, "Start", nil, nil), + wfTestNode("retrieve_1", workflowregistry.NodeTypeKnowledgeRetrieve, "Retrieve", wfTestInputs("query", "start_1", "userMessage"), config), + wfTestNode("end_1", workflowregistry.NodeTypeEnd, "End", nil, nil), + }, + []dsl.Edge{ + wfTestEdge("start_1", "retrieve_1", "edge_start_retrieve"), + wfTestEdge("retrieve_1", "end_1", "edge_retrieve_end"), + }, + ) +} + func ticketDraftReadyWorkflowDefinition() dsl.Definition { return wfTestDefinition( []dsl.Node{ diff --git a/internal/ai/workflow/registry/registry.go b/internal/ai/workflow/registry/registry.go index f9cb5dc..64144a3 100644 --- a/internal/ai/workflow/registry/registry.go +++ b/internal/ai/workflow/registry/registry.go @@ -32,7 +32,6 @@ func DefaultRegistry() *Registry { output("messageId", "消息 ID", VariableTypeInteger, "客户本轮消息的内部编号。"), output("aiAgentId", "AI Agent ID", VariableTypeInteger, "当前处理会话的 AI Agent 编号。"), output("userMessage", "用户消息", VariableTypeString, "客户本轮发送的原始消息内容。"), - output("knowledgeBaseIds", "知识库 ID 列表", VariableTypeIntegerArray, "当前 AI Agent 已绑定的知识库编号列表。"), }, }, NodeSpec{ @@ -120,6 +119,14 @@ func DefaultRegistry() *Registry { Description: "Retrieve knowledge for the current user message.", Icon: "BookOpenIcon", RiskLevel: NodeRiskLevelLow, + ConfigSchema: map[string]any{ + "knowledgeBaseIds": map[string]any{ + "type": string(VariableTypeIntegerArray), + "label": "知识库", + "required": true, + "description": "本节点检索时使用的知识库列表,按顺序表示优先级。", + }, + }, InputSchema: []VariableSpec{ requiredInput("query", "检索问题", VariableTypeString, "用于检索知识库的客户问题或查询文本。"), }, diff --git a/internal/ai/workflow/registry/registry_test.go b/internal/ai/workflow/registry/registry_test.go index 3ef41ac..539067f 100644 --- a/internal/ai/workflow/registry/registry_test.go +++ b/internal/ai/workflow/registry/registry_test.go @@ -23,8 +23,8 @@ func TestDefaultRegistryExposesStartOutputs(t *testing.T) { if !hasVariable(spec.OutputSchema, "userMessage", VariableTypeString) { t.Fatalf("expected start output userMessage:string, got %#v", spec.OutputSchema) } - if !hasVariable(spec.OutputSchema, "knowledgeBaseIds", VariableTypeIntegerArray) { - t.Fatalf("expected start output knowledgeBaseIds:array, got %#v", spec.OutputSchema) + if hasVariableName(spec.OutputSchema, "knowledgeBaseIds") { + t.Fatalf("did not expect start output knowledgeBaseIds after knowledge binding moved to retrieve node, got %#v", spec.OutputSchema) } } @@ -39,6 +39,9 @@ func TestDefaultRegistryExposesKnowledgeRetrieveInputsAndOutputs(t *testing.T) { if !hasVariable(spec.OutputSchema, "items", VariableTypeObjectArray) { t.Fatalf("expected knowledge_retrieve output items:array, got %#v", spec.OutputSchema) } + if spec.ConfigSchema == nil { + t.Fatalf("expected knowledge_retrieve config schema") + } } func TestDefaultRegistryExposesSendReplyRequiredInput(t *testing.T) { diff --git a/internal/ai/workflow/validator/validator.go b/internal/ai/workflow/validator/validator.go index 037782a..794072c 100644 --- a/internal/ai/workflow/validator/validator.go +++ b/internal/ai/workflow/validator/validator.go @@ -53,6 +53,7 @@ type definitionValidator struct { func (v *definitionValidator) validate() { v.validateNodes() v.validateEdges() + v.validateKnowledgeRetrieveConfigs() v.validateReachability() v.validateConfirmationGuards() v.validateVariableMappings() @@ -306,6 +307,58 @@ func (v *definitionValidator) validateConditions() { } } +func (v *definitionValidator) validateKnowledgeRetrieveConfigs() { + for index, node := range v.def.Nodes { + if strings.TrimSpace(node.Type) != registry.NodeTypeKnowledgeRetrieve { + continue + } + field := fmt.Sprintf("nodes[%d].config.knowledgeBaseIds", index) + ids, ok := readKnowledgeBaseIDsFromConfig(node.Data.Config) + if !ok || len(ids) == 0 { + v.addError(field, "知识检索节点需要选择至少一个知识库") + continue + } + for _, id := range ids { + if id <= 0 { + v.addError(field, "知识库 ID 必须大于 0") + break + } + } + } +} + +func readKnowledgeBaseIDsFromConfig(raw json.RawMessage) ([]int64, bool) { + if len(raw) == 0 { + return nil, false + } + var cfg map[string]any + if err := json.Unmarshal(raw, &cfg); err != nil { + return nil, false + } + rawIDs, ok := cfg["knowledgeBaseIds"] + if !ok { + return nil, false + } + values, ok := rawIDs.([]any) + if !ok { + return nil, false + } + ret := make([]int64, 0, len(values)) + for _, value := range values { + switch v := value.(type) { + case float64: + ret = append(ret, int64(v)) + case int64: + ret = append(ret, v) + case int: + ret = append(ret, int64(v)) + default: + return nil, false + } + } + return ret, true +} + func (v *definitionValidator) validateCondition(field string, sourceNodeID string, condition *dsl.Condition) { if condition == nil { v.addError(field, "condition branch condition is required") diff --git a/internal/ai/workflow/validator/validator_test.go b/internal/ai/workflow/validator/validator_test.go index 1dfdabb..18402b1 100644 --- a/internal/ai/workflow/validator/validator_test.go +++ b/internal/ai/workflow/validator/validator_test.go @@ -213,6 +213,42 @@ func TestValidateDefinitionRejectsUnknownConditionVariable(t *testing.T) { } } +func TestValidateDefinitionRejectsKnowledgeRetrieveWithoutKnowledgeBases(t *testing.T) { + def := knowledgeRetrieveDefinition(nil) + + result := validator.ValidateDefinition(def, registry.DefaultRegistry()) + + if result.Valid { + t.Fatalf("expected knowledge_retrieve without knowledge bases to be invalid") + } + if !hasValidationMessage(result, "需要选择至少一个知识库") { + t.Fatalf("expected missing knowledge base error, got %#v", result.Errors) + } +} + +func TestValidateDefinitionRejectsKnowledgeRetrieveWithInvalidKnowledgeBaseID(t *testing.T) { + def := knowledgeRetrieveDefinition(map[string]any{"knowledgeBaseIds": []int64{0, -1}}) + + result := validator.ValidateDefinition(def, registry.DefaultRegistry()) + + if result.Valid { + t.Fatalf("expected invalid knowledge base id to be invalid") + } + if !hasValidationMessage(result, "知识库 ID 必须大于 0") { + t.Fatalf("expected invalid knowledge base id error, got %#v", result.Errors) + } +} + +func TestValidateDefinitionAcceptsKnowledgeRetrieveWithKnowledgeBases(t *testing.T) { + def := knowledgeRetrieveDefinition(map[string]any{"knowledgeBaseIds": []int64{1, 2}}) + + result := validator.ValidateDefinition(def, registry.DefaultRegistry()) + + if !result.Valid { + t.Fatalf("expected knowledge_retrieve with knowledge bases to be valid, got %#v", result.Errors) + } +} + func minimalDefinition() dsl.Definition { return dsl.Definition{ SchemaVersion: dsl.SchemaVersion, @@ -228,6 +264,21 @@ func minimalDefinition() dsl.Definition { } } +func knowledgeRetrieveDefinition(config any) dsl.Definition { + return dsl.Definition{ + SchemaVersion: dsl.SchemaVersion, + Nodes: []dsl.Node{ + node("start_1", "start", nil, nil), + node("retrieve_1", "knowledge_retrieve", inputs("query", dsl.RefValue("start_1", "userMessage")), config), + node("end_1", "end", nil, nil), + }, + Edges: []dsl.Edge{ + edge("start_1", "retrieve_1"), + edge("retrieve_1", "end_1"), + }, + } +} + func conditionDefinition() dsl.Definition { conditionConfig := dsl.ConditionConfig{ Branches: []dsl.ConditionBranch{ diff --git a/internal/handlers/dashboard/ai_agent_handler.go b/internal/handlers/dashboard/ai_agent_handler.go index a81adc4..1fc159a 100644 --- a/internal/handlers/dashboard/ai_agent_handler.go +++ b/internal/handlers/dashboard/ai_agent_handler.go @@ -182,9 +182,7 @@ func buildAIAgentResponseWithLocale(item *models.AIAgent, locale string) respons FallbackMode: item.FallbackMode, FallbackModeName: enums.GetAIAgentFallbackModeLabel(item.FallbackMode), FallbackMessage: item.FallbackMessage, - KnowledgeIDs: utils.SplitInt64s(item.KnowledgeIDs), SkillIDs: utils.SplitInt64s(item.SkillIDs), - KnowledgeBaseNames: make([]string, 0), Skills: make([]response.AIAgentSkillResponse, 0), Teams: make([]response.AIAgentTeamResponse, 0), DirectTools: make([]response.AIAgentMCPToolResponse, 0), @@ -209,11 +207,6 @@ func buildAIAgentResponseWithLocale(item *models.AIAgent, locale string) respons }) } } - for _, id := range ret.KnowledgeIDs { - if knowledgeBase := services.KnowledgeBaseService.Get(id); knowledgeBase != nil { - ret.KnowledgeBaseNames = append(ret.KnowledgeBaseNames, knowledgeBase.Name) - } - } for _, id := range ret.SkillIDs { if skill := services.SkillDefinitionService.Get(id); skill != nil { ret.Skills = append(ret.Skills, response.AIAgentSkillResponse{ diff --git a/internal/pkg/dto/request/ai_request.go b/internal/pkg/dto/request/ai_request.go index 457f35d..286e3ce 100644 --- a/internal/pkg/dto/request/ai_request.go +++ b/internal/pkg/dto/request/ai_request.go @@ -54,7 +54,6 @@ type CreateAIAgentRequest struct { HandoffMode enums.AIAgentHandoffMode `json:"handoffMode"` FallbackMode enums.AIAgentFallbackMode `json:"fallbackMode"` FallbackMessage string `json:"fallbackMessage"` - KnowledgeIDs []int64 `json:"knowledgeIds"` SkillIDs []int64 `json:"skillIds"` DirectTools []AIAgentMCPToolRequest `json:"directTools"` } diff --git a/internal/pkg/dto/response/ai_response.go b/internal/pkg/dto/response/ai_response.go index 0ab19f3..ffa5c8e 100644 --- a/internal/pkg/dto/response/ai_response.go +++ b/internal/pkg/dto/response/ai_response.go @@ -85,8 +85,6 @@ type AIAgentResponse struct { FallbackMode enums.AIAgentFallbackMode `json:"fallbackMode"` FallbackModeName string `json:"fallbackModeName"` FallbackMessage string `json:"fallbackMessage"` - KnowledgeIDs []int64 `json:"knowledgeIds"` - KnowledgeBaseNames []string `json:"knowledgeBaseNames"` SkillIDs []int64 `json:"skillIds"` Skills []AIAgentSkillResponse `json:"skills"` DirectTools []AIAgentMCPToolResponse `json:"directTools"` diff --git a/internal/repositories/ai_agent_repository.go b/internal/repositories/ai_agent_repository.go index 165b047..f56992f 100644 --- a/internal/repositories/ai_agent_repository.go +++ b/internal/repositories/ai_agent_repository.go @@ -1,11 +1,7 @@ package repositories import ( - "strconv" - "agent-desk/internal/models" - "agent-desk/internal/pkg/enums" - "agent-desk/internal/pkg/httpx/params" "github.com/mlogclub/simple/sqls" @@ -112,16 +108,3 @@ func (r *aIAgentRepository) FindByIds(db *gorm.DB, ids []int64) []models.AIAgent db.Where("id IN ?", ids).Find(&list) return list } - -func (r *aIAgentRepository) FindByKnowledgeBaseID(db *gorm.DB, knowledgeBaseID int64) (list []models.AIAgent) { - id := strconv.FormatInt(knowledgeBaseID, 10) - db.Where( - "(knowledge_ids = ? OR knowledge_ids LIKE ? OR knowledge_ids LIKE ? OR knowledge_ids LIKE ?) AND status <> ?", - id, - id+",%", - "%,"+id, - "%,"+id+",%", - enums.StatusDeleted, - ).Find(&list) - return -} diff --git a/internal/services/ai_agent_service.go b/internal/services/ai_agent_service.go index d91929c..21bdd04 100644 --- a/internal/services/ai_agent_service.go +++ b/internal/services/ai_agent_service.go @@ -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{}) diff --git a/internal/services/ai_agent_workflow_service_test.go b/internal/services/ai_agent_workflow_service_test.go index d909c8d..347500f 100644 --- a/internal/services/ai_agent_workflow_service_test.go +++ b/internal/services/ai_agent_workflow_service_test.go @@ -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 { diff --git a/internal/services/ai_workflow_service.go b/internal/services/ai_workflow_service.go index e47a024..204e63f 100644 --- a/internal/services/ai_workflow_service.go +++ b/internal/services/ai_workflow_service.go @@ -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"), diff --git a/internal/services/dashboard_service.go b/internal/services/dashboard_service.go index 9fd30d1..27a5f0d 100644 --- a/internal/services/dashboard_service.go +++ b/internal/services/dashboard_service.go @@ -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 diff --git a/internal/services/knowledge_base_service.go b/internal/services/knowledge_base_service.go index d269fc9..c6a3043 100644 --- a/internal/services/knowledge_base_service.go +++ b/internal/services/knowledge_base_service.go @@ -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 { diff --git a/internal/services/knowledge_base_service_test.go b/internal/services/knowledge_base_service_test.go index aa6b7b9..7e5410d 100644 --- a/internal/services/knowledge_base_service_test.go +++ b/internal/services/knowledge_base_service_test.go @@ -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{ diff --git a/web/app/dashboard/ai-agents/_components/config-workbench.tsx b/web/app/dashboard/ai-agents/_components/config-workbench.tsx index 52f68b4..6a0e0c0 100644 --- a/web/app/dashboard/ai-agents/_components/config-workbench.tsx +++ b/web/app/dashboard/ai-agents/_components/config-workbench.tsx @@ -2,8 +2,6 @@ import { useCallback, useEffect, useMemo, useState, type ReactNode } from "react" import { - ArrowDownIcon, - ArrowUpIcon, BotMessageSquareIcon, GitBranchIcon, HistoryIcon, @@ -44,7 +42,6 @@ import { fetchAIWorkflowNodeSpecs, fetchAIWorkflowVersions, fetchAgentTeamsAll, - fetchKnowledgeBasesAll, fetchMCPCatalog, fetchSkillDefinitionsAll, publishAIAgentWorkflow, @@ -58,7 +55,6 @@ import { type AIWorkflowVersion, type AdminAgentTeam, type CreateAIAgentPayload, - type KnowledgeBase, type MCPToolCatalogItem, type MCPToolSourceType, type SkillDefinition, @@ -148,7 +144,6 @@ export function AIAgentConfigWorkbench({ const [handoffMode, setHandoffMode] = useState(String(AIAgentHandoffMode.WaitPool)) const [fallbackMode, setFallbackMode] = useState(String(AIAgentFallbackMode.NoAnswer)) const [fallbackMessage, setFallbackMessage] = useState("") - const [selectedKnowledgeIds, setSelectedKnowledgeIds] = useState([]) const [selectedTeamIds, setSelectedTeamIds] = useState([]) const [selectedSkillIds, setSelectedSkillIds] = useState([]) const [directTools, setDirectTools] = useState([]) @@ -157,11 +152,9 @@ export function AIAgentConfigWorkbench({ const [workflowRevision, setWorkflowRevision] = useState(0) const [aiConfigs, setAIConfigs] = useState([]) - const [knowledgeBases, setKnowledgeBases] = useState([]) const [agentTeams, setAgentTeams] = useState([]) const [skills, setSkills] = useState([]) const [toolCatalog, setToolCatalog] = useState([]) - const [knowledgeToAdd, setKnowledgeToAdd] = useState("") const [teamToAdd, setTeamToAdd] = useState("") const [skillToAdd, setSkillToAdd] = useState("") const [directToolGroupToAdd, setDirectToolGroupToAdd] = useState("") @@ -183,7 +176,6 @@ export function AIAgentConfigWorkbench({ specs, defaultDefinition, configs, - bases, teams, skillList, catalog, @@ -191,7 +183,6 @@ export function AIAgentConfigWorkbench({ fetchAIWorkflowNodeSpecs(), fetchAIWorkflowDefaultDefinition().catch(() => fallbackDefinition), fetchAIConfigsAll({ modelType: AIModelType.LLM }), - fetchKnowledgeBasesAll({ status: Status.Ok }), fetchAgentTeamsAll(), fetchSkillDefinitionsAll({ status: Status.Ok }), fetchMCPCatalog(), @@ -199,7 +190,6 @@ export function AIAgentConfigWorkbench({ setNodeSpecs(specs ?? []) setAIConfigs(configs ?? []) - setKnowledgeBases(bases ?? []) setAgentTeams(teams ?? []) setSkills(skillList ?? []) setToolCatalog(catalog ?? []) @@ -217,7 +207,6 @@ export function AIAgentConfigWorkbench({ setHandoffMode(String(AIAgentHandoffMode.WaitPool)) setFallbackMode(String(AIAgentFallbackMode.NoAnswer)) setFallbackMessage("") - setSelectedKnowledgeIds([]) setSelectedTeamIds([]) setSelectedSkillIds([]) setDirectTools([]) @@ -247,7 +236,6 @@ export function AIAgentConfigWorkbench({ setHandoffMode(String(agentDetail.handoffMode || AIAgentHandoffMode.WaitPool)) setFallbackMode(String(agentDetail.fallbackMode || AIAgentFallbackMode.NoAnswer)) setFallbackMessage(agentDetail.fallbackMessage || "") - setSelectedKnowledgeIds(agentDetail.knowledgeIds ?? []) setSelectedTeamIds((agentDetail.teams ?? []).map((team) => team.id)) setSelectedSkillIds(agentDetail.skillIds ?? []) setDirectTools(agentDetail.directTools ?? []) @@ -290,10 +278,6 @@ export function AIAgentConfigWorkbench({ () => aiConfigs.map((item) => ({ value: String(item.id), label: `${item.name} · ${item.modelName}` })), [aiConfigs] ) - const knowledgeOptions = useMemo( - () => knowledgeBases.map((item) => ({ value: String(item.id), label: item.name })), - [knowledgeBases] - ) const teamOptions = useMemo( () => agentTeams.map((item) => ({ value: String(item.id), label: item.name })), [agentTeams] @@ -356,16 +340,6 @@ export function AIAgentConfigWorkbench({ setNext([...current, id]) } - function moveKnowledge(index: number, direction: -1 | 1) { - const targetIndex = index + direction - if (targetIndex < 0 || targetIndex >= selectedKnowledgeIds.length) return - const next = [...selectedKnowledgeIds] - const current = next[index] - next[index] = next[targetIndex] - next[targetIndex] = current - setSelectedKnowledgeIds(next) - } - function addDirectTool(value: string) { const option = directToolOptions.find((item) => item.value === value) if (!option) return @@ -390,7 +364,6 @@ export function AIAgentConfigWorkbench({ handoffMode: Number(handoffMode), fallbackMode: Number(fallbackMode), fallbackMessage: fallbackMessage.trim(), - knowledgeIds: uniqueNumbers(selectedKnowledgeIds), skillIds: uniqueNumbers(selectedSkillIds), directTools, } @@ -505,7 +478,6 @@ export function AIAgentConfigWorkbench({ { key: "workflow", title: "会话流程", icon: }, ] - const selectedKnowledgeOptions = selectedOptions(selectedKnowledgeIds, knowledgeOptions) const selectedTeamOptions = selectedOptions(selectedTeamIds, teamOptions) const selectedSkillOptions = selectedOptions(selectedSkillIds, skillOptions) const workflowPublished = isWorkflowPublished(agent) @@ -683,59 +655,6 @@ export function AIAgentConfigWorkbench({ ) : null} - {activeSection === "capabilities" ? ( - - !selectedKnowledgeIds.includes(Number(option.value)))} - placeholder="选择知识库" - onValueChange={setKnowledgeToAdd} - onAdd={() => { - addSelected(knowledgeToAdd, selectedKnowledgeIds, setSelectedKnowledgeIds) - setKnowledgeToAdd("") - }} - /> -
- {selectedKnowledgeOptions.length === 0 ? ( -
至少选择一个知识库。
- ) : ( - selectedKnowledgeOptions.map((option, index) => ( -
- {index + 1} -
{option.label}
- - - -
- )) - )} -
-
- ) : null} - {activeSection === "capabilities" ? ( { - const knowledgeIds = item.knowledgeIds ?? []; - const knowledgeBaseNames = item.knowledgeBaseNames ?? []; - return ( -
- {knowledgeIds.length === 0 ? ( - - {t("aiAgent.notConfigured")} - - ) : ( - knowledgeBaseNames.map((name, index) => ( - - {name} - - )) - )} -
- ); - }, - }, { key: "skills", label: t("aiAgent.columnSkills"), diff --git a/web/app/dashboard/ai-workflows/_components/node-config-panel.tsx b/web/app/dashboard/ai-workflows/_components/node-config-panel.tsx index 27fbd8b..da0e0c4 100644 --- a/web/app/dashboard/ai-workflows/_components/node-config-panel.tsx +++ b/web/app/dashboard/ai-workflows/_components/node-config-panel.tsx @@ -1,13 +1,14 @@ "use client" -import type { ReactNode } from "react" -import { Trash2Icon } from "lucide-react" +import { useEffect, useMemo, useState, type ReactNode } from "react" +import { ArrowDownIcon, ArrowUpIcon, Trash2Icon } from "lucide-react" import { Button } from "@/components/ui/button" import { Input } from "@/components/ui/input" import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs" import { OptionCombobox } from "@/components/option-combobox" -import type { AIWorkflowDefinition, AIWorkflowNodeSpec } from "@/lib/api/admin" +import { fetchKnowledgeBasesAll, type AIWorkflowDefinition, type AIWorkflowNodeSpec, type KnowledgeBase } from "@/lib/api/admin" +import { Status } from "@/lib/generated/enums" import { cn } from "@/lib/utils" import { VariableSelector } from "./variable-selector" @@ -177,6 +178,13 @@ export function NodeConfigPanel({ outputContent={outputFields} /> + {node.type === "knowledge_retrieve" ? ( + updateConfig(nextConfig)} + /> + ) : null} + {showConditionBranches && (node.type === "condition" || branches.length > 0) ? ( + onChange: (config: Record) => void +}) { + const [knowledgeBases, setKnowledgeBases] = useState([]) + const [knowledgeToAdd, setKnowledgeToAdd] = useState("") + const selectedKnowledgeIds = normalizeKnowledgeBaseIds(config.knowledgeBaseIds) + const knowledgeOptions = useMemo( + () => knowledgeBases.map((item) => ({ value: String(item.id), label: item.name })), + [knowledgeBases] + ) + const selectedKnowledgeOptions = selectedKnowledgeIds + .map((id) => knowledgeOptions.find((option) => Number(option.value) === id)) + .filter((option): option is { value: string; label: string } => Boolean(option)) + + useEffect(() => { + let cancelled = false + fetchKnowledgeBasesAll({ status: Status.Ok }) + .then((items) => { + if (!cancelled) { + setKnowledgeBases(items ?? []) + } + }) + .catch(() => { + if (!cancelled) { + setKnowledgeBases([]) + } + }) + return () => { + cancelled = true + } + }, []) + + const updateKnowledgeBaseIds = (ids: number[]) => { + onChange({ ...config, knowledgeBaseIds: uniquePositiveNumbers(ids) }) + } + const addKnowledgeBase = () => { + const id = Number(knowledgeToAdd) + if (!Number.isFinite(id) || id <= 0 || selectedKnowledgeIds.includes(id)) return + updateKnowledgeBaseIds([...selectedKnowledgeIds, id]) + setKnowledgeToAdd("") + } + const moveKnowledgeBase = (index: number, direction: -1 | 1) => { + const targetIndex = index + direction + if (targetIndex < 0 || targetIndex >= selectedKnowledgeIds.length) return + const next = [...selectedKnowledgeIds] + const current = next[index] + next[index] = next[targetIndex] + next[targetIndex] = current + updateKnowledgeBaseIds(next) + } + + return ( + + +
+
+ !selectedKnowledgeIds.includes(Number(option.value)))} + placeholder="选择知识库" + searchPlaceholder="搜索知识库" + emptyText="没有可用知识库" + triggerClassName={inspectorComboboxClassName} + onChange={setKnowledgeToAdd} + /> + +
+ +
+ {selectedKnowledgeOptions.length === 0 ? ( +
+ 未选择知识库,流程发布校验不会通过。 +
+ ) : ( + selectedKnowledgeOptions.map((option, index) => ( +
+ + {index + 1} + +
{option.label}
+
+ + + +
+
+ )) + )} +
+
+
+
+ ) +} + function ConditionBranchesEditor({ branches, nodes, @@ -639,3 +777,18 @@ function stringifyConditionRight(value: unknown) { } return JSON.stringify(value) } + +function normalizeKnowledgeBaseIds(value: unknown) { + if (!Array.isArray(value)) { + return [] + } + return uniquePositiveNumbers( + value + .map((item) => Number(item)) + .filter((item) => Number.isFinite(item)) + ) +} + +function uniquePositiveNumbers(input: number[]) { + return Array.from(new Set(input.filter((item) => item > 0))) +} diff --git a/web/app/dashboard/ai-workflows/_components/workflow-utils.test.mjs b/web/app/dashboard/ai-workflows/_components/workflow-utils.test.mjs index 7a1ffc7..09dbdd2 100644 --- a/web/app/dashboard/ai-workflows/_components/workflow-utils.test.mjs +++ b/web/app/dashboard/ai-workflows/_components/workflow-utils.test.mjs @@ -140,6 +140,27 @@ describe("validateWorkflowDefinition", () => { assert.deepEqual(plain(result), { valid: true, errors: [] }) }) + + it("rejects knowledge retrieve nodes without node knowledge bases", async () => { + const { createRefValue, validateWorkflowDefinition } = await loadModule() + + const result = validateWorkflowDefinition({ + schemaVersion: 2, + nodes: [ + workflowNode("start_1", "start"), + workflowNode("retrieve_1", "knowledge_retrieve", { x: 240, y: 0 }, { + title: "知识检索", + inputsValues: { query: createRefValue("start_1", "userMessage") }, + config: { knowledgeBaseIds: [] }, + }), + workflowNode("end_1", "end", { x: 480, y: 0 }), + ], + edges: [workflowEdge("start_1", "retrieve_1"), workflowEdge("retrieve_1", "end_1")], + }) + + assert.equal(result.valid, false) + assert.match(result.errors.join("\n"), /需要选择至少一个知识库/) + }) }) describe("createWorkflowNodeFromSpec", () => { diff --git a/web/app/dashboard/ai-workflows/_components/workflow-utils.ts b/web/app/dashboard/ai-workflows/_components/workflow-utils.ts index bffa294..8a25cab 100644 --- a/web/app/dashboard/ai-workflows/_components/workflow-utils.ts +++ b/web/app/dashboard/ai-workflows/_components/workflow-utils.ts @@ -179,6 +179,15 @@ export function validateWorkflowDefinition( if (!node.type?.trim()) { errors.push(`node type is required: ${node.id}`) } + if (node.type === "knowledge_retrieve") { + const config = normalizeNodeConfig(node.data?.config) + const knowledgeBaseIds = Array.isArray(config.knowledgeBaseIds) ? config.knowledgeBaseIds : [] + if (knowledgeBaseIds.length === 0) { + errors.push(`${getNodeTitle(node, nodeSpecs)} 需要选择至少一个知识库`) + } else if (knowledgeBaseIds.some((id) => Number(id) <= 0)) { + errors.push(`${getNodeTitle(node, nodeSpecs)} 知识库 ID 必须大于 0`) + } + } } for (const edge of edges) { diff --git a/web/lib/api/admin.ts b/web/lib/api/admin.ts index 658c774..4dfd256 100644 --- a/web/lib/api/admin.ts +++ b/web/lib/api/admin.ts @@ -227,8 +227,6 @@ export type AIAgent = { fallbackMode: number fallbackModeName: string fallbackMessage: string - knowledgeIds: number[] - knowledgeBaseNames: string[] skillIds: number[] skills: { id: number; name: string }[] directTools: { @@ -262,7 +260,6 @@ export type CreateAIAgentPayload = { handoffMode: number fallbackMode: number fallbackMessage: string - knowledgeIds: number[] skillIds: number[] directTools: { toolCode: string