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/ai/runtime/traces"
"agent-desk/internal/models" "agent-desk/internal/models"
"agent-desk/internal/pkg/enums" "agent-desk/internal/pkg/enums"
"agent-desk/internal/pkg/utils"
"agent-desk/internal/repositories" "agent-desk/internal/repositories"
"github.com/mlogclub/simple/sqls" "github.com/mlogclub/simple/sqls"
@@ -21,6 +20,7 @@ const defaultRuntimeKnowledgeMaxContextItems = 5
type KnowledgeRetriever struct { type KnowledgeRetriever struct {
AIAgent models.AIAgent AIAgent models.AIAgent
knowledgeBaseIDs []int64
} }
type KnowledgeRetrieveOptions struct { type KnowledgeRetrieveOptions struct {
@@ -52,8 +52,11 @@ type KnowledgeRetrieveResult struct {
Policies []KnowledgeBaseRetrievePolicy Policies []KnowledgeBaseRetrievePolicy
} }
func NewKnowledgeRetriever(aiAgent models.AIAgent) *KnowledgeRetriever { func NewKnowledgeRetriever(aiAgent models.AIAgent, knowledgeBaseIDs []int64) *KnowledgeRetriever {
return &KnowledgeRetriever{AIAgent: aiAgent} return &KnowledgeRetriever{
AIAgent: aiAgent,
knowledgeBaseIDs: append([]int64(nil), knowledgeBaseIDs...),
}
} }
func DefaultKnowledgeRetrieveOptions() KnowledgeRetrieveOptions { func DefaultKnowledgeRetrieveOptions() KnowledgeRetrieveOptions {
@@ -63,8 +66,8 @@ func DefaultKnowledgeRetrieveOptions() KnowledgeRetrieveOptions {
} }
} }
func (r *KnowledgeRetriever) KnowledgeBaseIDs() []int64 { func (r *KnowledgeRetriever) ConfiguredKnowledgeBaseIDs() []int64 {
return utils.SplitInt64s(r.AIAgent.KnowledgeIDs) return append([]int64(nil), r.knowledgeBaseIDs...)
} }
func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) { 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) { 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{ return rag.Retrieve.RetrieveWithTrace(ctx, rag.RetrieveRequest{
Query: query, Query: query,
KnowledgeBaseIDs: ids, 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) { func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) (*KnowledgeRetrieveResult, error) {
query = strings.TrimSpace(query) query = strings.TrimSpace(query)
knowledgeBaseIDs := r.KnowledgeBaseIDs() knowledgeBaseIDs := r.ConfiguredKnowledgeBaseIDs()
policies := r.resolvePolicies(knowledgeBaseIDs, opts) policies := r.resolvePolicies(knowledgeBaseIDs, opts)
contextMaxTokens := opts.ContextMaxTokens contextMaxTokens := opts.ContextMaxTokens
if contextMaxTokens <= 0 { if contextMaxTokens <= 0 {
+32 -4
View File
@@ -20,7 +20,6 @@ import (
"agent-desk/internal/pkg/dto" "agent-desk/internal/pkg/dto"
"agent-desk/internal/pkg/dto/request" "agent-desk/internal/pkg/dto/request"
"agent-desk/internal/pkg/enums" "agent-desk/internal/pkg/enums"
"agent-desk/internal/pkg/utils"
"agent-desk/internal/services" "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, "messageId": state.input.UserMessage.ID,
"aiAgentId": state.input.AIAgent.ID, "aiAgentId": state.input.AIAgent.ID,
"userMessage": strings.TrimSpace(state.input.UserMessage.Content), "userMessage": strings.TrimSpace(state.input.UserMessage.Content),
"knowledgeBaseIds": utils.SplitInt64s(state.input.AIAgent.KnowledgeIDs),
"conversationState": state.input.Conversation.Status, "conversationState": state.input.Conversation.Status,
}) })
case workflowregistry.NodeTypeConversationUnderstanding: 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 { func (e *Executor) executeKnowledgeRetrieve(ctx context.Context, state *runState, node dsl.Node) error {
query := strings.TrimSpace(toString(state.resolveInput(node, "query"))) 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) result, err := retriever.RetrieveContext(ctx, query)
if err != nil { if err != nil {
return err 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 != "" { if prompt := strings.TrimSpace(readStringConfig(node.Data.Config, "prompt")); prompt != "" {
systemPrompt = strings.TrimSpace(systemPrompt + "\n\n" + 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)}) state.setNodeVars(node.ID, map[string]any{"replyText": workflowKnowledgeFallbackReply(state.input.AIAgent)})
return nil return nil
} }
@@ -1088,6 +1090,32 @@ func readBoolConfig(raw json.RawMessage, key string) bool {
return truthy(cfg[key]) 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 { func compareString(left any, right any) int {
return strings.Compare(toString(left), toString(right)) return strings.Compare(toString(left), toString(right))
} }
+32 -3
View File
@@ -298,7 +298,6 @@ func TestExecutorPolicyFirstWorkflowRoutesGreetingToDirectReply(t *testing.T) {
Content: "<p>你好。</p>", Content: "<p>你好。</p>",
}, },
AIAgent: models.AIAgent{ AIAgent: models.AIAgent{
KnowledgeIDs: "1",
FallbackMessage: "我暂时没有找到足够准确的信息。", FallbackMessage: "我暂时没有找到足够准确的信息。",
}, },
}) })
@@ -330,7 +329,6 @@ func TestExecutorPolicyFirstWorkflowRoutesBusinessQuestionToKnowledge(t *testing
Content: "你们价格是多少?", Content: "你们价格是多少?",
}, },
AIAgent: models.AIAgent{ AIAgent: models.AIAgent{
KnowledgeIDs: "1",
FallbackMessage: "我暂时没有找到足够准确的信息。", 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"}) 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) { func TestExecutorLLMReplyUsesAgentFallbackWhenDeclaredKnowledgeIsEmpty(t *testing.T) {
result, err := NewExecutor().Execute(context.Background(), Input{ result, err := NewExecutor().Execute(context.Background(), Input{
Definition: emptyKnowledgeReplyDefinition(), Definition: emptyKnowledgeReplyDefinition(),
@@ -347,7 +363,6 @@ func TestExecutorLLMReplyUsesAgentFallbackWhenDeclaredKnowledgeIsEmpty(t *testin
Content: "产品功能", Content: "产品功能",
}, },
AIAgent: models.AIAgent{ AIAgent: models.AIAgent{
KnowledgeIDs: "1",
FallbackMode: enums.AIAgentFallbackModeNoAnswer, FallbackMode: enums.AIAgentFallbackModeNoAnswer,
FallbackMessage: "我暂时没有找到足够准确的信息。你可以补充更具体的问题,我再继续帮你查。", FallbackMessage: "我暂时没有找到足够准确的信息。你可以补充更具体的问题,我再继续帮你查。",
SystemPrompt: "不要编造事实。", 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 { func ticketDraftReadyWorkflowDefinition() dsl.Definition {
return wfTestDefinition( return wfTestDefinition(
[]dsl.Node{ []dsl.Node{
+8 -1
View File
@@ -32,7 +32,6 @@ func DefaultRegistry() *Registry {
output("messageId", "消息 ID", VariableTypeInteger, "客户本轮消息的内部编号。"), output("messageId", "消息 ID", VariableTypeInteger, "客户本轮消息的内部编号。"),
output("aiAgentId", "AI Agent ID", VariableTypeInteger, "当前处理会话的 AI Agent 编号。"), output("aiAgentId", "AI Agent ID", VariableTypeInteger, "当前处理会话的 AI Agent 编号。"),
output("userMessage", "用户消息", VariableTypeString, "客户本轮发送的原始消息内容。"), output("userMessage", "用户消息", VariableTypeString, "客户本轮发送的原始消息内容。"),
output("knowledgeBaseIds", "知识库 ID 列表", VariableTypeIntegerArray, "当前 AI Agent 已绑定的知识库编号列表。"),
}, },
}, },
NodeSpec{ NodeSpec{
@@ -120,6 +119,14 @@ func DefaultRegistry() *Registry {
Description: "Retrieve knowledge for the current user message.", Description: "Retrieve knowledge for the current user message.",
Icon: "BookOpenIcon", Icon: "BookOpenIcon",
RiskLevel: NodeRiskLevelLow, RiskLevel: NodeRiskLevelLow,
ConfigSchema: map[string]any{
"knowledgeBaseIds": map[string]any{
"type": string(VariableTypeIntegerArray),
"label": "知识库",
"required": true,
"description": "本节点检索时使用的知识库列表,按顺序表示优先级。",
},
},
InputSchema: []VariableSpec{ InputSchema: []VariableSpec{
requiredInput("query", "检索问题", VariableTypeString, "用于检索知识库的客户问题或查询文本。"), requiredInput("query", "检索问题", VariableTypeString, "用于检索知识库的客户问题或查询文本。"),
}, },
@@ -23,8 +23,8 @@ func TestDefaultRegistryExposesStartOutputs(t *testing.T) {
if !hasVariable(spec.OutputSchema, "userMessage", VariableTypeString) { if !hasVariable(spec.OutputSchema, "userMessage", VariableTypeString) {
t.Fatalf("expected start output userMessage:string, got %#v", spec.OutputSchema) t.Fatalf("expected start output userMessage:string, got %#v", spec.OutputSchema)
} }
if !hasVariable(spec.OutputSchema, "knowledgeBaseIds", VariableTypeIntegerArray) { if hasVariableName(spec.OutputSchema, "knowledgeBaseIds") {
t.Fatalf("expected start output knowledgeBaseIds:array<int>, got %#v", spec.OutputSchema) 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) { if !hasVariable(spec.OutputSchema, "items", VariableTypeObjectArray) {
t.Fatalf("expected knowledge_retrieve output items:array<object>, got %#v", spec.OutputSchema) 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) { func TestDefaultRegistryExposesSendReplyRequiredInput(t *testing.T) {
@@ -53,6 +53,7 @@ type definitionValidator struct {
func (v *definitionValidator) validate() { func (v *definitionValidator) validate() {
v.validateNodes() v.validateNodes()
v.validateEdges() v.validateEdges()
v.validateKnowledgeRetrieveConfigs()
v.validateReachability() v.validateReachability()
v.validateConfirmationGuards() v.validateConfirmationGuards()
v.validateVariableMappings() 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) { func (v *definitionValidator) validateCondition(field string, sourceNodeID string, condition *dsl.Condition) {
if condition == nil { if condition == nil {
v.addError(field, "condition branch condition is required") 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 { func minimalDefinition() dsl.Definition {
return dsl.Definition{ return dsl.Definition{
SchemaVersion: dsl.SchemaVersion, 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 { func conditionDefinition() dsl.Definition {
conditionConfig := dsl.ConditionConfig{ conditionConfig := dsl.ConditionConfig{
Branches: []dsl.ConditionBranch{ Branches: []dsl.ConditionBranch{
@@ -182,9 +182,7 @@ func buildAIAgentResponseWithLocale(item *models.AIAgent, locale string) respons
FallbackMode: item.FallbackMode, FallbackMode: item.FallbackMode,
FallbackModeName: enums.GetAIAgentFallbackModeLabel(item.FallbackMode), FallbackModeName: enums.GetAIAgentFallbackModeLabel(item.FallbackMode),
FallbackMessage: item.FallbackMessage, FallbackMessage: item.FallbackMessage,
KnowledgeIDs: utils.SplitInt64s(item.KnowledgeIDs),
SkillIDs: utils.SplitInt64s(item.SkillIDs), SkillIDs: utils.SplitInt64s(item.SkillIDs),
KnowledgeBaseNames: make([]string, 0),
Skills: make([]response.AIAgentSkillResponse, 0), Skills: make([]response.AIAgentSkillResponse, 0),
Teams: make([]response.AIAgentTeamResponse, 0), Teams: make([]response.AIAgentTeamResponse, 0),
DirectTools: make([]response.AIAgentMCPToolResponse, 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 { for _, id := range ret.SkillIDs {
if skill := services.SkillDefinitionService.Get(id); skill != nil { if skill := services.SkillDefinitionService.Get(id); skill != nil {
ret.Skills = append(ret.Skills, response.AIAgentSkillResponse{ ret.Skills = append(ret.Skills, response.AIAgentSkillResponse{
-1
View File
@@ -54,7 +54,6 @@ type CreateAIAgentRequest struct {
HandoffMode enums.AIAgentHandoffMode `json:"handoffMode"` HandoffMode enums.AIAgentHandoffMode `json:"handoffMode"`
FallbackMode enums.AIAgentFallbackMode `json:"fallbackMode"` FallbackMode enums.AIAgentFallbackMode `json:"fallbackMode"`
FallbackMessage string `json:"fallbackMessage"` FallbackMessage string `json:"fallbackMessage"`
KnowledgeIDs []int64 `json:"knowledgeIds"`
SkillIDs []int64 `json:"skillIds"` SkillIDs []int64 `json:"skillIds"`
DirectTools []AIAgentMCPToolRequest `json:"directTools"` DirectTools []AIAgentMCPToolRequest `json:"directTools"`
} }
-2
View File
@@ -85,8 +85,6 @@ type AIAgentResponse struct {
FallbackMode enums.AIAgentFallbackMode `json:"fallbackMode"` FallbackMode enums.AIAgentFallbackMode `json:"fallbackMode"`
FallbackModeName string `json:"fallbackModeName"` FallbackModeName string `json:"fallbackModeName"`
FallbackMessage string `json:"fallbackMessage"` FallbackMessage string `json:"fallbackMessage"`
KnowledgeIDs []int64 `json:"knowledgeIds"`
KnowledgeBaseNames []string `json:"knowledgeBaseNames"`
SkillIDs []int64 `json:"skillIds"` SkillIDs []int64 `json:"skillIds"`
Skills []AIAgentSkillResponse `json:"skills"` Skills []AIAgentSkillResponse `json:"skills"`
DirectTools []AIAgentMCPToolResponse `json:"directTools"` DirectTools []AIAgentMCPToolResponse `json:"directTools"`
@@ -1,11 +1,7 @@
package repositories package repositories
import ( import (
"strconv"
"agent-desk/internal/models" "agent-desk/internal/models"
"agent-desk/internal/pkg/enums"
"agent-desk/internal/pkg/httpx/params" "agent-desk/internal/pkg/httpx/params"
"github.com/mlogclub/simple/sqls" "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) db.Where("id IN ?", ids).Find(&list)
return 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, "handoff_mode": item.HandoffMode,
"fallback_mode": item.FallbackMode, "fallback_mode": item.FallbackMode,
"fallback_message": item.FallbackMessage, "fallback_message": item.FallbackMessage,
"knowledge_ids": item.KnowledgeIDs,
"skill_ids": item.SkillIDs, "skill_ids": item.SkillIDs,
"allowed_mcp_tools": item.AllowedMCPTools, "allowed_mcp_tools": item.AllowedMCPTools,
"update_user_id": operator.UserID, "update_user_id": operator.UserID,
@@ -177,13 +176,6 @@ func (s *aIAgentService) buildAIAgentModel(id int64, req request.CreateAIAgentRe
return nil, errorsx.InvalidParamI18n("error.e0144") 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) skillIDs, err := s.normalizeSkillIDs(req.SkillIDs)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -212,7 +204,6 @@ func (s *aIAgentService) buildAIAgentModel(id int64, req request.CreateAIAgentRe
HandoffMode: req.HandoffMode, HandoffMode: req.HandoffMode,
FallbackMode: req.FallbackMode, FallbackMode: req.FallbackMode,
FallbackMessage: strings.TrimSpace(req.FallbackMessage), FallbackMessage: strings.TrimSpace(req.FallbackMessage),
KnowledgeIDs: utils.JoinInt64s(knowledgeIDs),
SkillIDs: utils.JoinInt64s(skillIDs), SkillIDs: utils.JoinInt64s(skillIDs),
AllowedMCPTools: directToolsJSON, AllowedMCPTools: directToolsJSON,
WorkflowVersionID: 0, WorkflowVersionID: 0,
@@ -243,29 +234,6 @@ func (s *aIAgentService) normalizeTeamIDs(input []int64) ([]int64, error) {
return ret, nil 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) { func (s *aIAgentService) normalizeSkillIDs(input []int64) ([]int64, error) {
ret := make([]int64, 0, len(input)) ret := make([]int64, 0, len(input))
seen := make(map[int64]struct{}) seen := make(map[int64]struct{})
@@ -22,7 +22,6 @@ func TestAIAgentServiceCreatesDefaultWorkflow(t *testing.T) {
setupAIAgentWorkflowTestDB(t) setupAIAgentWorkflowTestDB(t)
operator := aiAgentWorkflowTestOperator() operator := aiAgentWorkflowTestOperator()
aiConfigID := createAIAgentWorkflowTestConfig(t) aiConfigID := createAIAgentWorkflowTestConfig(t)
knowledgeID := createAIAgentWorkflowTestKnowledgeBase(t)
item, err := AIAgentService.CreateAIAgent(request.CreateAIAgentRequest{ item, err := AIAgentService.CreateAIAgent(request.CreateAIAgentRequest{
Name: "workflow agent", Name: "workflow agent",
@@ -30,7 +29,6 @@ func TestAIAgentServiceCreatesDefaultWorkflow(t *testing.T) {
ServiceMode: enums.IMConversationServiceModeAIOnly, ServiceMode: enums.IMConversationServiceModeAIOnly,
HandoffMode: enums.AIAgentHandoffModeWaitPool, HandoffMode: enums.AIAgentHandoffModeWaitPool,
FallbackMode: enums.AIAgentFallbackModeNoAnswer, FallbackMode: enums.AIAgentFallbackModeNoAnswer,
KnowledgeIDs: []int64{knowledgeID},
}, operator) }, operator)
if err != nil { if err != nil {
t.Fatalf("CreateAIAgent() error = %v", err) t.Fatalf("CreateAIAgent() error = %v", err)
@@ -54,8 +52,8 @@ func TestAIAgentServiceCreatesDefaultWorkflow(t *testing.T) {
t.Fatalf("expected default draft definition") t.Fatalf("expected default draft definition")
} }
validation := workflowvalidator.ValidateDefinition(stored, workflowregistry.DefaultRegistry()) validation := workflowvalidator.ValidateDefinition(stored, workflowregistry.DefaultRegistry())
if !validation.Valid { if validation.Valid || !workflowValidationHasMessage(validation, "需要选择至少一个知识库") {
t.Fatalf("expected default workflow to be valid, got %#v", validation.Errors) t.Fatalf("expected default workflow to require node knowledge bases, got %#v", validation.Errors)
} }
if nodeTypeByID(stored, "understanding_1") != workflowregistry.NodeTypeConversationUnderstanding { if nodeTypeByID(stored, "understanding_1") != workflowregistry.NodeTypeConversationUnderstanding {
t.Fatalf("expected default workflow to include conversation understanding, got nodes: %#v", stored.Nodes) 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() definition := AIWorkflowService.DefaultAgentWorkflowDefinition()
if definition.SchemaVersion != dsl.SchemaVersion || nodeTypeByID(definition, "start_1") != workflowregistry.NodeTypeStart { if definition.SchemaVersion != dsl.SchemaVersion || nodeTypeByID(definition, "start_1") != workflowregistry.NodeTypeStart {
t.Fatalf("expected default workflow definition") t.Fatalf("expected default workflow definition")
} }
validation := workflowvalidator.ValidateDefinition(definition, workflowregistry.DefaultRegistry()) validation := workflowvalidator.ValidateDefinition(definition, workflowregistry.DefaultRegistry())
if !validation.Valid { if validation.Valid || !workflowValidationHasMessage(validation, "需要选择至少一个知识库") {
t.Fatalf("expected default workflow definition to be valid, got %#v", validation.Errors) t.Fatalf("expected default workflow definition to require node knowledge bases, got %#v", validation.Errors)
} }
if nodeTypeByID(definition, "understanding_1") != workflowregistry.NodeTypeConversationUnderstanding { if nodeTypeByID(definition, "understanding_1") != workflowregistry.NodeTypeConversationUnderstanding {
t.Fatalf("expected default workflow to include conversation understanding, got nodes: %#v", definition.Nodes) 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) setupAIAgentWorkflowTestDB(t)
operator := aiAgentWorkflowTestOperator() operator := aiAgentWorkflowTestOperator()
aiConfigID := createAIAgentWorkflowTestConfig(t) aiConfigID := createAIAgentWorkflowTestConfig(t)
knowledgeID := createAIAgentWorkflowTestKnowledgeBase(t)
agent, err := AIAgentService.CreateAIAgent(request.CreateAIAgentRequest{ agent, err := AIAgentService.CreateAIAgent(request.CreateAIAgentRequest{
Name: "workflow agent without version", Name: "workflow agent without version",
@@ -175,7 +172,6 @@ func TestAIWorkflowServicePublishAgentWorkflowBindsAgentVersion(t *testing.T) {
ServiceMode: enums.IMConversationServiceModeAIOnly, ServiceMode: enums.IMConversationServiceModeAIOnly,
HandoffMode: enums.AIAgentHandoffModeWaitPool, HandoffMode: enums.AIAgentHandoffModeWaitPool,
FallbackMode: enums.AIAgentFallbackModeNoAnswer, FallbackMode: enums.AIAgentFallbackModeNoAnswer,
KnowledgeIDs: []int64{knowledgeID},
}, operator) }, operator)
if err != nil { if err != nil {
t.Fatalf("CreateAIAgent() error = %v", err) t.Fatalf("CreateAIAgent() error = %v", err)
@@ -287,6 +283,15 @@ func workflowHasNodeType(def dsl.Definition, nodeType string) bool {
return false 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 { func nodeTypeByID(def dsl.Definition, nodeID string) string {
for _, node := range def.Nodes { for _, node := range def.Nodes {
if node.ID == nodeID { if node.ID == nodeID {
+1 -1
View File
@@ -472,7 +472,7 @@ func defaultAgentWorkflowDefinition() dsl.Definition {
"followUpQuestions": dsl.RefValue("draft_ticket_1", "followUpQuestions"), "followUpQuestions": dsl.RefValue("draft_ticket_1", "followUpQuestions"),
}, map[string]any{"staticReply": "为了创建工单,还需要补充以下信息:\n{{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("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{ workflowNode("answerability_1", workflowregistry.NodeTypeAnswerabilityGate, "可回答判断", 2940, 753, map[string]dsl.Value{
"userMessage": dsl.RefValue("start_1", "userMessage"), "userMessage": dsl.RefValue("start_1", "userMessage"),
"knowledgeItems": dsl.RefValue("retrieve_1", "items"), "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 { sort.Slice(alerts, func(i, j int) bool {
if alerts[i].Count == alerts[j].Count { if alerts[i].Count == alerts[j].Count {
return alerts[i].ID < alerts[j].ID return alerts[i].ID < alerts[j].ID
+90 -5
View File
@@ -2,9 +2,14 @@ package services
import ( import (
"context" "context"
"encoding/json"
"fmt"
"strings"
"time" "time"
"agent-desk/internal/ai/rag" "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/models"
"agent-desk/internal/pkg/dto" "agent-desk/internal/pkg/dto"
"agent-desk/internal/pkg/dto/request" "agent-desk/internal/pkg/dto/request"
@@ -128,12 +133,12 @@ func (s *knowledgeBaseService) DeleteKnowledgeBase(id int64) error {
return errorsx.InvalidParamI18n("error.e0283") return errorsx.InvalidParamI18n("error.e0283")
} }
referencingAgents := repositories.AIAgentRepository.FindByKnowledgeBaseID(sqls.DB(), id) referencingWorkflows := s.findWorkflowReferencesByKnowledgeBaseID(id)
if len(referencingAgents) > 0 { if len(referencingWorkflows) > 0 {
if len(referencingAgents) == 1 { if len(referencingWorkflows) == 1 {
return errorsx.ForbiddenI18n("error.knowledgeBase.referencedByAgent", referencingAgents[0].Name) 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 { 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) 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 { func (s *knowledgeBaseService) UpdateSort(ids []int64) error {
return sqls.WithTransaction(func(ctx *sqls.TxContext) error { return sqls.WithTransaction(func(ctx *sqls.TxContext) error {
for i, id := range ids { for i, id := range ids {
@@ -1,10 +1,12 @@
package services package services
import ( import (
"fmt" "encoding/json"
"strings" "strings"
"testing" "testing"
"agent-desk/internal/ai/workflow/dsl"
workflowregistry "agent-desk/internal/ai/workflow/registry"
"agent-desk/internal/models" "agent-desk/internal/models"
"agent-desk/internal/pkg/dto/request" "agent-desk/internal/pkg/dto/request"
"agent-desk/internal/pkg/enums" "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) setupKnowledgeBaseServiceTestDB(t)
kb := createKnowledgeBaseServiceTestBase(t, "Referenced KB") kb := createKnowledgeBaseServiceTestBase(t, "Referenced KB")
otherKB := createKnowledgeBaseServiceTestBase(t, "Other KB") otherKB := createKnowledgeBaseServiceTestBase(t, "Other KB")
if err := repositories.AIAgentRepository.Create(sqls.DB(), &models.AIAgent{ createKnowledgeBaseServiceTestWorkflow(t, "Support Workflow", knowledgeBaseServiceTestWorkflowDefinition([]int64{12, otherKB.ID}))
Name: "Support Agent", createKnowledgeBaseServiceTestWorkflow(t, "Knowledge Workflow", knowledgeBaseServiceTestWorkflowDefinition([]int64{12, kb.ID, otherKB.ID}))
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)
}
err := KnowledgeBaseService.DeleteKnowledgeBase(kb.ID) err := KnowledgeBaseService.DeleteKnowledgeBase(kb.ID)
if err == nil { 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") { if got := err.Error(); !strings.Contains(got, "Knowledge Workflow") {
t.Fatalf("DeleteKnowledgeBase() error = %q, want agent name", got) t.Fatalf("DeleteKnowledgeBase() error = %q, want workflow name", got)
} }
if repositories.KnowledgeBaseRepository.Get(sqls.DB(), kb.ID) == nil { 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 { if err != nil {
t.Fatalf("open sqlite db: %v", err) 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) t.Fatalf("auto migrate: %v", err)
} }
sqls.SetDB(db) 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 { func createKnowledgeBaseServiceTestBase(t *testing.T, name string) *models.KnowledgeBase {
t.Helper() t.Helper()
item := &models.KnowledgeBase{ item := &models.KnowledgeBase{
@@ -2,8 +2,6 @@
import { useCallback, useEffect, useMemo, useState, type ReactNode } from "react" import { useCallback, useEffect, useMemo, useState, type ReactNode } from "react"
import { import {
ArrowDownIcon,
ArrowUpIcon,
BotMessageSquareIcon, BotMessageSquareIcon,
GitBranchIcon, GitBranchIcon,
HistoryIcon, HistoryIcon,
@@ -44,7 +42,6 @@ import {
fetchAIWorkflowNodeSpecs, fetchAIWorkflowNodeSpecs,
fetchAIWorkflowVersions, fetchAIWorkflowVersions,
fetchAgentTeamsAll, fetchAgentTeamsAll,
fetchKnowledgeBasesAll,
fetchMCPCatalog, fetchMCPCatalog,
fetchSkillDefinitionsAll, fetchSkillDefinitionsAll,
publishAIAgentWorkflow, publishAIAgentWorkflow,
@@ -58,7 +55,6 @@ import {
type AIWorkflowVersion, type AIWorkflowVersion,
type AdminAgentTeam, type AdminAgentTeam,
type CreateAIAgentPayload, type CreateAIAgentPayload,
type KnowledgeBase,
type MCPToolCatalogItem, type MCPToolCatalogItem,
type MCPToolSourceType, type MCPToolSourceType,
type SkillDefinition, type SkillDefinition,
@@ -148,7 +144,6 @@ export function AIAgentConfigWorkbench({
const [handoffMode, setHandoffMode] = useState(String(AIAgentHandoffMode.WaitPool)) const [handoffMode, setHandoffMode] = useState(String(AIAgentHandoffMode.WaitPool))
const [fallbackMode, setFallbackMode] = useState(String(AIAgentFallbackMode.NoAnswer)) const [fallbackMode, setFallbackMode] = useState(String(AIAgentFallbackMode.NoAnswer))
const [fallbackMessage, setFallbackMessage] = useState("") const [fallbackMessage, setFallbackMessage] = useState("")
const [selectedKnowledgeIds, setSelectedKnowledgeIds] = useState<number[]>([])
const [selectedTeamIds, setSelectedTeamIds] = useState<number[]>([]) const [selectedTeamIds, setSelectedTeamIds] = useState<number[]>([])
const [selectedSkillIds, setSelectedSkillIds] = useState<number[]>([]) const [selectedSkillIds, setSelectedSkillIds] = useState<number[]>([])
const [directTools, setDirectTools] = useState<DirectToolItem[]>([]) const [directTools, setDirectTools] = useState<DirectToolItem[]>([])
@@ -157,11 +152,9 @@ export function AIAgentConfigWorkbench({
const [workflowRevision, setWorkflowRevision] = useState(0) const [workflowRevision, setWorkflowRevision] = useState(0)
const [aiConfigs, setAIConfigs] = useState<AIConfig[]>([]) const [aiConfigs, setAIConfigs] = useState<AIConfig[]>([])
const [knowledgeBases, setKnowledgeBases] = useState<KnowledgeBase[]>([])
const [agentTeams, setAgentTeams] = useState<AdminAgentTeam[]>([]) const [agentTeams, setAgentTeams] = useState<AdminAgentTeam[]>([])
const [skills, setSkills] = useState<SkillDefinition[]>([]) const [skills, setSkills] = useState<SkillDefinition[]>([])
const [toolCatalog, setToolCatalog] = useState<MCPToolCatalogItem[]>([]) const [toolCatalog, setToolCatalog] = useState<MCPToolCatalogItem[]>([])
const [knowledgeToAdd, setKnowledgeToAdd] = useState("")
const [teamToAdd, setTeamToAdd] = useState("") const [teamToAdd, setTeamToAdd] = useState("")
const [skillToAdd, setSkillToAdd] = useState("") const [skillToAdd, setSkillToAdd] = useState("")
const [directToolGroupToAdd, setDirectToolGroupToAdd] = useState("") const [directToolGroupToAdd, setDirectToolGroupToAdd] = useState("")
@@ -183,7 +176,6 @@ export function AIAgentConfigWorkbench({
specs, specs,
defaultDefinition, defaultDefinition,
configs, configs,
bases,
teams, teams,
skillList, skillList,
catalog, catalog,
@@ -191,7 +183,6 @@ export function AIAgentConfigWorkbench({
fetchAIWorkflowNodeSpecs(), fetchAIWorkflowNodeSpecs(),
fetchAIWorkflowDefaultDefinition().catch(() => fallbackDefinition), fetchAIWorkflowDefaultDefinition().catch(() => fallbackDefinition),
fetchAIConfigsAll({ modelType: AIModelType.LLM }), fetchAIConfigsAll({ modelType: AIModelType.LLM }),
fetchKnowledgeBasesAll({ status: Status.Ok }),
fetchAgentTeamsAll(), fetchAgentTeamsAll(),
fetchSkillDefinitionsAll({ status: Status.Ok }), fetchSkillDefinitionsAll({ status: Status.Ok }),
fetchMCPCatalog(), fetchMCPCatalog(),
@@ -199,7 +190,6 @@ export function AIAgentConfigWorkbench({
setNodeSpecs(specs ?? []) setNodeSpecs(specs ?? [])
setAIConfigs(configs ?? []) setAIConfigs(configs ?? [])
setKnowledgeBases(bases ?? [])
setAgentTeams(teams ?? []) setAgentTeams(teams ?? [])
setSkills(skillList ?? []) setSkills(skillList ?? [])
setToolCatalog(catalog ?? []) setToolCatalog(catalog ?? [])
@@ -217,7 +207,6 @@ export function AIAgentConfigWorkbench({
setHandoffMode(String(AIAgentHandoffMode.WaitPool)) setHandoffMode(String(AIAgentHandoffMode.WaitPool))
setFallbackMode(String(AIAgentFallbackMode.NoAnswer)) setFallbackMode(String(AIAgentFallbackMode.NoAnswer))
setFallbackMessage("") setFallbackMessage("")
setSelectedKnowledgeIds([])
setSelectedTeamIds([]) setSelectedTeamIds([])
setSelectedSkillIds([]) setSelectedSkillIds([])
setDirectTools([]) setDirectTools([])
@@ -247,7 +236,6 @@ export function AIAgentConfigWorkbench({
setHandoffMode(String(agentDetail.handoffMode || AIAgentHandoffMode.WaitPool)) setHandoffMode(String(agentDetail.handoffMode || AIAgentHandoffMode.WaitPool))
setFallbackMode(String(agentDetail.fallbackMode || AIAgentFallbackMode.NoAnswer)) setFallbackMode(String(agentDetail.fallbackMode || AIAgentFallbackMode.NoAnswer))
setFallbackMessage(agentDetail.fallbackMessage || "") setFallbackMessage(agentDetail.fallbackMessage || "")
setSelectedKnowledgeIds(agentDetail.knowledgeIds ?? [])
setSelectedTeamIds((agentDetail.teams ?? []).map((team) => team.id)) setSelectedTeamIds((agentDetail.teams ?? []).map((team) => team.id))
setSelectedSkillIds(agentDetail.skillIds ?? []) setSelectedSkillIds(agentDetail.skillIds ?? [])
setDirectTools(agentDetail.directTools ?? []) setDirectTools(agentDetail.directTools ?? [])
@@ -290,10 +278,6 @@ export function AIAgentConfigWorkbench({
() => aiConfigs.map((item) => ({ value: String(item.id), label: `${item.name} · ${item.modelName}` })), () => aiConfigs.map((item) => ({ value: String(item.id), label: `${item.name} · ${item.modelName}` })),
[aiConfigs] [aiConfigs]
) )
const knowledgeOptions = useMemo(
() => knowledgeBases.map((item) => ({ value: String(item.id), label: item.name })),
[knowledgeBases]
)
const teamOptions = useMemo( const teamOptions = useMemo(
() => agentTeams.map((item) => ({ value: String(item.id), label: item.name })), () => agentTeams.map((item) => ({ value: String(item.id), label: item.name })),
[agentTeams] [agentTeams]
@@ -356,16 +340,6 @@ export function AIAgentConfigWorkbench({
setNext([...current, id]) 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) { function addDirectTool(value: string) {
const option = directToolOptions.find((item) => item.value === value) const option = directToolOptions.find((item) => item.value === value)
if (!option) return if (!option) return
@@ -390,7 +364,6 @@ export function AIAgentConfigWorkbench({
handoffMode: Number(handoffMode), handoffMode: Number(handoffMode),
fallbackMode: Number(fallbackMode), fallbackMode: Number(fallbackMode),
fallbackMessage: fallbackMessage.trim(), fallbackMessage: fallbackMessage.trim(),
knowledgeIds: uniqueNumbers(selectedKnowledgeIds),
skillIds: uniqueNumbers(selectedSkillIds), skillIds: uniqueNumbers(selectedSkillIds),
directTools, directTools,
} }
@@ -505,7 +478,6 @@ export function AIAgentConfigWorkbench({
{ key: "workflow", title: "会话流程", icon: <GitBranchIcon /> }, { key: "workflow", title: "会话流程", icon: <GitBranchIcon /> },
] ]
const selectedKnowledgeOptions = selectedOptions(selectedKnowledgeIds, knowledgeOptions)
const selectedTeamOptions = selectedOptions(selectedTeamIds, teamOptions) const selectedTeamOptions = selectedOptions(selectedTeamIds, teamOptions)
const selectedSkillOptions = selectedOptions(selectedSkillIds, skillOptions) const selectedSkillOptions = selectedOptions(selectedSkillIds, skillOptions)
const workflowPublished = isWorkflowPublished(agent) const workflowPublished = isWorkflowPublished(agent)
@@ -683,59 +655,6 @@ export function AIAgentConfigWorkbench({
</ConfigSection> </ConfigSection>
) : null} ) : 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" ? ( {activeSection === "capabilities" ? (
<ConfigSection> <ConfigSection>
<AddRow <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", key: "skills",
label: t("aiAgent.columnSkills"), label: t("aiAgent.columnSkills"),
@@ -1,13 +1,14 @@
"use client" "use client"
import type { ReactNode } from "react" import { useEffect, useMemo, useState, type ReactNode } from "react"
import { Trash2Icon } from "lucide-react" import { ArrowDownIcon, ArrowUpIcon, Trash2Icon } from "lucide-react"
import { Button } from "@/components/ui/button" import { Button } from "@/components/ui/button"
import { Input } from "@/components/ui/input" import { Input } from "@/components/ui/input"
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs" import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"
import { OptionCombobox } from "@/components/option-combobox" 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 { cn } from "@/lib/utils"
import { VariableSelector } from "./variable-selector" import { VariableSelector } from "./variable-selector"
@@ -177,6 +178,13 @@ export function NodeConfigPanel({
outputContent={outputFields} outputContent={outputFields}
/> />
{node.type === "knowledge_retrieve" ? (
<KnowledgeRetrieveConfigPanel
config={config}
onChange={(nextConfig) => updateConfig(nextConfig)}
/>
) : null}
{showConditionBranches && (node.type === "condition" || branches.length > 0) ? ( {showConditionBranches && (node.type === "condition" || branches.length > 0) ? (
<ConditionBranchesEditor <ConditionBranchesEditor
branches={branches} 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({ function ConditionBranchesEditor({
branches, branches,
nodes, nodes,
@@ -639,3 +777,18 @@ function stringifyConditionRight(value: unknown) {
} }
return JSON.stringify(value) 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: [] }) 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", () => { describe("createWorkflowNodeFromSpec", () => {
@@ -179,6 +179,15 @@ export function validateWorkflowDefinition(
if (!node.type?.trim()) { if (!node.type?.trim()) {
errors.push(`node type is required: ${node.id}`) 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) { for (const edge of edges) {
-3
View File
@@ -227,8 +227,6 @@ export type AIAgent = {
fallbackMode: number fallbackMode: number
fallbackModeName: string fallbackModeName: string
fallbackMessage: string fallbackMessage: string
knowledgeIds: number[]
knowledgeBaseNames: string[]
skillIds: number[] skillIds: number[]
skills: { id: number; name: string }[] skills: { id: number; name: string }[]
directTools: { directTools: {
@@ -262,7 +260,6 @@ export type CreateAIAgentPayload = {
handoffMode: number handoffMode: number
fallbackMode: number fallbackMode: number
fallbackMessage: string fallbackMessage: string
knowledgeIds: number[]
skillIds: number[] skillIds: number[]
directTools: { directTools: {
toolCode: string toolCode: string