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:
@@ -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 {
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user