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
@@ -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 {
+32 -4
View File
@@ -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))
}
+32 -3
View File
@@ -298,7 +298,6 @@ func TestExecutorPolicyFirstWorkflowRoutesGreetingToDirectReply(t *testing.T) {
Content: "<p>你好。</p>",
},
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{
+8 -1
View File
@@ -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, "用于检索知识库的客户问题或查询文本。"),
},
@@ -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<int>, 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<object>, got %#v", spec.OutputSchema)
}
if spec.ConfigSchema == nil {
t.Fatalf("expected knowledge_retrieve config schema")
}
}
func TestDefaultRegistryExposesSendReplyRequiredInput(t *testing.T) {
@@ -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")
@@ -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{
@@ -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{
-1
View File
@@ -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"`
}
-2
View File
@@ -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"`
@@ -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
}
-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{
@@ -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<number[]>([])
const [selectedTeamIds, setSelectedTeamIds] = useState<number[]>([])
const [selectedSkillIds, setSelectedSkillIds] = useState<number[]>([])
const [directTools, setDirectTools] = useState<DirectToolItem[]>([])
@@ -157,11 +152,9 @@ export function AIAgentConfigWorkbench({
const [workflowRevision, setWorkflowRevision] = useState(0)
const [aiConfigs, setAIConfigs] = useState<AIConfig[]>([])
const [knowledgeBases, setKnowledgeBases] = useState<KnowledgeBase[]>([])
const [agentTeams, setAgentTeams] = useState<AdminAgentTeam[]>([])
const [skills, setSkills] = useState<SkillDefinition[]>([])
const [toolCatalog, setToolCatalog] = useState<MCPToolCatalogItem[]>([])
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: <GitBranchIcon /> },
]
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({
</ConfigSection>
) : null}
{activeSection === "capabilities" ? (
<ConfigSection>
<AddRow
value={knowledgeToAdd}
options={knowledgeOptions.filter((option) => !selectedKnowledgeIds.includes(Number(option.value)))}
placeholder="选择知识库"
onValueChange={setKnowledgeToAdd}
onAdd={() => {
addSelected(knowledgeToAdd, selectedKnowledgeIds, setSelectedKnowledgeIds)
setKnowledgeToAdd("")
}}
/>
<div className="space-y-2 rounded-md border p-3">
{selectedKnowledgeOptions.length === 0 ? (
<div className="text-sm text-muted-foreground"></div>
) : (
selectedKnowledgeOptions.map((option, index) => (
<div key={option.value} className="flex items-center gap-2">
<Badge variant="secondary" className="min-w-8 justify-center">{index + 1}</Badge>
<div className="flex-1 text-sm">{option.label}</div>
<Button
type="button"
variant="outline"
size="icon-sm"
disabled={index === 0}
onClick={() => moveKnowledge(index, -1)}
>
<ArrowUpIcon />
</Button>
<Button
type="button"
variant="outline"
size="icon-sm"
disabled={index === selectedKnowledgeOptions.length - 1}
onClick={() => moveKnowledge(index, 1)}
>
<ArrowDownIcon />
</Button>
<Button
type="button"
variant="outline"
size="icon-sm"
onClick={() => setSelectedKnowledgeIds((current) => current.filter((id) => id !== Number(option.value)))}
>
<Trash2Icon />
</Button>
</div>
))
)}
</div>
</ConfigSection>
) : null}
{activeSection === "capabilities" ? (
<ConfigSection>
<AddRow
-26
View File
@@ -150,32 +150,6 @@ export default function DashboardAIAgentsPage() {
);
},
},
{
key: "knowledge",
label: t("aiAgent.columnKnowledge"),
render: (item) => {
const knowledgeIds = item.knowledgeIds ?? [];
const knowledgeBaseNames = item.knowledgeBaseNames ?? [];
return (
<div className="flex flex-wrap gap-1">
{knowledgeIds.length === 0 ? (
<span className="text-sm text-muted-foreground">
{t("aiAgent.notConfigured")}
</span>
) : (
knowledgeBaseNames.map((name, index) => (
<Badge
key={knowledgeIds[index] ?? `${item.id}-${index}`}
variant="secondary"
>
{name}
</Badge>
))
)}
</div>
);
},
},
{
key: "skills",
label: t("aiAgent.columnSkills"),
@@ -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" ? (
<KnowledgeRetrieveConfigPanel
config={config}
onChange={(nextConfig) => updateConfig(nextConfig)}
/>
) : null}
{showConditionBranches && (node.type === "condition" || branches.length > 0) ? (
<ConditionBranchesEditor
branches={branches}
@@ -277,6 +285,136 @@ export function ConditionBranchConfigPanel({
)
}
function KnowledgeRetrieveConfigPanel({
config,
onChange,
}: {
config: Record<string, unknown>
onChange: (config: Record<string, unknown>) => void
}) {
const [knowledgeBases, setKnowledgeBases] = useState<KnowledgeBase[]>([])
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 (
<InspectorSection title="节点配置" meta={`${selectedKnowledgeIds.length} 个知识库`}>
<InspectorRow label="知识库" required>
<div className="space-y-2">
<div className="grid grid-cols-[minmax(0,1fr)_auto] gap-2">
<OptionCombobox
value={knowledgeToAdd}
options={knowledgeOptions.filter((option) => !selectedKnowledgeIds.includes(Number(option.value)))}
placeholder="选择知识库"
searchPlaceholder="搜索知识库"
emptyText="没有可用知识库"
triggerClassName={inspectorComboboxClassName}
onChange={setKnowledgeToAdd}
/>
<Button type="button" variant="outline" size="sm" className="h-8" onClick={addKnowledgeBase}>
</Button>
</div>
<div className="divide-y divide-slate-100 rounded-md border border-slate-200">
{selectedKnowledgeOptions.length === 0 ? (
<div className="px-3 py-2 text-sm text-slate-500">
</div>
) : (
selectedKnowledgeOptions.map((option, index) => (
<div key={option.value} className="grid grid-cols-[32px_minmax(0,1fr)_auto] items-center gap-2 px-2 py-1.5">
<span className="inline-flex h-5 items-center justify-center rounded-sm border border-slate-200 bg-slate-50 font-mono text-xs text-slate-500">
{index + 1}
</span>
<div className="truncate text-sm text-slate-700">{option.label}</div>
<div className="flex items-center gap-1">
<Button
type="button"
variant="ghost"
size="icon"
className="size-7"
disabled={index === 0}
onClick={() => moveKnowledgeBase(index, -1)}
aria-label="上移知识库"
>
<ArrowUpIcon className="size-3.5" />
</Button>
<Button
type="button"
variant="ghost"
size="icon"
className="size-7"
disabled={index === selectedKnowledgeOptions.length - 1}
onClick={() => moveKnowledgeBase(index, 1)}
aria-label="下移知识库"
>
<ArrowDownIcon className="size-3.5" />
</Button>
<Button
type="button"
variant="ghost"
size="icon"
className="size-7 text-slate-500 hover:text-destructive"
onClick={() => updateKnowledgeBaseIds(selectedKnowledgeIds.filter((id) => id !== Number(option.value)))}
aria-label="移除知识库"
>
<Trash2Icon className="size-3.5" />
</Button>
</div>
</div>
))
)}
</div>
</div>
</InspectorRow>
</InspectorSection>
)
}
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)))
}
@@ -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", () => {
@@ -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) {
-3
View File
@@ -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