diff --git a/internal/ai/runtime/retrievers/knowledge_retriever.go b/internal/ai/runtime/retrievers/knowledge_retriever.go
index 13704fb..661799b 100644
--- a/internal/ai/runtime/retrievers/knowledge_retriever.go
+++ b/internal/ai/runtime/retrievers/knowledge_retriever.go
@@ -8,7 +8,6 @@ import (
"agent-desk/internal/ai/runtime/traces"
"agent-desk/internal/models"
"agent-desk/internal/pkg/enums"
- "agent-desk/internal/pkg/utils"
"agent-desk/internal/repositories"
"github.com/mlogclub/simple/sqls"
@@ -20,7 +19,8 @@ const defaultRuntimeKnowledgeScoreThreshold = 0.3
const defaultRuntimeKnowledgeMaxContextItems = 5
type KnowledgeRetriever struct {
- AIAgent models.AIAgent
+ AIAgent models.AIAgent
+ knowledgeBaseIDs []int64
}
type KnowledgeRetrieveOptions struct {
@@ -52,8 +52,11 @@ type KnowledgeRetrieveResult struct {
Policies []KnowledgeBaseRetrievePolicy
}
-func NewKnowledgeRetriever(aiAgent models.AIAgent) *KnowledgeRetriever {
- return &KnowledgeRetriever{AIAgent: aiAgent}
+func NewKnowledgeRetriever(aiAgent models.AIAgent, knowledgeBaseIDs []int64) *KnowledgeRetriever {
+ return &KnowledgeRetriever{
+ AIAgent: aiAgent,
+ knowledgeBaseIDs: append([]int64(nil), knowledgeBaseIDs...),
+ }
}
func DefaultKnowledgeRetrieveOptions() KnowledgeRetrieveOptions {
@@ -63,8 +66,8 @@ func DefaultKnowledgeRetrieveOptions() KnowledgeRetrieveOptions {
}
}
-func (r *KnowledgeRetriever) KnowledgeBaseIDs() []int64 {
- return utils.SplitInt64s(r.AIAgent.KnowledgeIDs)
+func (r *KnowledgeRetriever) ConfiguredKnowledgeBaseIDs() []int64 {
+ return append([]int64(nil), r.knowledgeBaseIDs...)
}
func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) {
@@ -72,7 +75,7 @@ func (r *KnowledgeRetriever) Retrieve(ctx context.Context, query string) ([]rag.
}
func (r *KnowledgeRetriever) RetrieveByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) ([]rag.RetrieveResult, *rag.RetrieveTrace, error) {
- ids := r.KnowledgeBaseIDs()
+ ids := r.ConfiguredKnowledgeBaseIDs()
return rag.Retrieve.RetrieveWithTrace(ctx, rag.RetrieveRequest{
Query: query,
KnowledgeBaseIDs: ids,
@@ -87,7 +90,7 @@ func (r *KnowledgeRetriever) RetrieveContext(ctx context.Context, query string)
func (r *KnowledgeRetriever) RetrieveContextByOptions(ctx context.Context, opts KnowledgeRetrieveOptions, query string) (*KnowledgeRetrieveResult, error) {
query = strings.TrimSpace(query)
- knowledgeBaseIDs := r.KnowledgeBaseIDs()
+ knowledgeBaseIDs := r.ConfiguredKnowledgeBaseIDs()
policies := r.resolvePolicies(knowledgeBaseIDs, opts)
contextMaxTokens := opts.ContextMaxTokens
if contextMaxTokens <= 0 {
diff --git a/internal/ai/runtime/workflow/executor.go b/internal/ai/runtime/workflow/executor.go
index 6b014e1..2d2e73b 100644
--- a/internal/ai/runtime/workflow/executor.go
+++ b/internal/ai/runtime/workflow/executor.go
@@ -20,7 +20,6 @@ import (
"agent-desk/internal/pkg/dto"
"agent-desk/internal/pkg/dto/request"
"agent-desk/internal/pkg/enums"
- "agent-desk/internal/pkg/utils"
"agent-desk/internal/services"
)
@@ -270,7 +269,6 @@ func (e *Executor) executeNode(ctx context.Context, state *runState, node dsl.No
"messageId": state.input.UserMessage.ID,
"aiAgentId": state.input.AIAgent.ID,
"userMessage": strings.TrimSpace(state.input.UserMessage.Content),
- "knowledgeBaseIds": utils.SplitInt64s(state.input.AIAgent.KnowledgeIDs),
"conversationState": state.input.Conversation.Status,
})
case workflowregistry.NodeTypeConversationUnderstanding:
@@ -732,7 +730,11 @@ func (e *Executor) executeHandoffToHuman(state *runState, node dsl.Node) error {
func (e *Executor) executeKnowledgeRetrieve(ctx context.Context, state *runState, node dsl.Node) error {
query := strings.TrimSpace(toString(state.resolveInput(node, "query")))
- retriever := retrievers.NewKnowledgeRetriever(state.input.AIAgent)
+ knowledgeBaseIDs := readInt64ArrayConfig(node.Data.Config, "knowledgeBaseIds")
+ if len(knowledgeBaseIDs) == 0 {
+ return fmt.Errorf("knowledge retrieve node requires knowledgeBaseIds")
+ }
+ retriever := retrievers.NewKnowledgeRetriever(state.input.AIAgent, knowledgeBaseIDs)
result, err := retriever.RetrieveContext(ctx, query)
if err != nil {
return err
@@ -784,7 +786,7 @@ func (e *Executor) executeLLMReply(ctx context.Context, state *runState, node ds
if prompt := strings.TrimSpace(readStringConfig(node.Data.Config, "prompt")); prompt != "" {
systemPrompt = strings.TrimSpace(systemPrompt + "\n\n" + prompt)
}
- if _, declaresKnowledge := node.Data.InputsValues["knowledgeItems"]; declaresKnowledge && len(utils.SplitInt64s(state.input.AIAgent.KnowledgeIDs)) > 0 && !hasItems(state.resolveInput(node, "knowledgeItems")) {
+ if _, declaresKnowledge := node.Data.InputsValues["knowledgeItems"]; declaresKnowledge && !hasItems(state.resolveInput(node, "knowledgeItems")) {
state.setNodeVars(node.ID, map[string]any{"replyText": workflowKnowledgeFallbackReply(state.input.AIAgent)})
return nil
}
@@ -1088,6 +1090,32 @@ func readBoolConfig(raw json.RawMessage, key string) bool {
return truthy(cfg[key])
}
+func readInt64ArrayConfig(raw json.RawMessage, key string) []int64 {
+ if len(raw) == 0 {
+ return nil
+ }
+ var cfg map[string]any
+ if err := json.Unmarshal(raw, &cfg); err != nil {
+ return nil
+ }
+ items, ok := cfg[key].([]any)
+ if !ok {
+ return nil
+ }
+ ret := make([]int64, 0, len(items))
+ for _, item := range items {
+ switch value := item.(type) {
+ case float64:
+ ret = append(ret, int64(value))
+ case int64:
+ ret = append(ret, value)
+ case int:
+ ret = append(ret, int64(value))
+ }
+ }
+ return ret
+}
+
func compareString(left any, right any) int {
return strings.Compare(toString(left), toString(right))
}
diff --git a/internal/ai/runtime/workflow/executor_test.go b/internal/ai/runtime/workflow/executor_test.go
index f6adc8f..36cd80d 100644
--- a/internal/ai/runtime/workflow/executor_test.go
+++ b/internal/ai/runtime/workflow/executor_test.go
@@ -298,7 +298,6 @@ func TestExecutorPolicyFirstWorkflowRoutesGreetingToDirectReply(t *testing.T) {
Content: "
你好。
",
},
AIAgent: models.AIAgent{
- KnowledgeIDs: "1",
FallbackMessage: "我暂时没有找到足够准确的信息。",
},
})
@@ -330,7 +329,6 @@ func TestExecutorPolicyFirstWorkflowRoutesBusinessQuestionToKnowledge(t *testing
Content: "你们价格是多少?",
},
AIAgent: models.AIAgent{
- KnowledgeIDs: "1",
FallbackMessage: "我暂时没有找到足够准确的信息。",
},
})
@@ -340,6 +338,24 @@ func TestExecutorPolicyFirstWorkflowRoutesBusinessQuestionToKnowledge(t *testing
assertPath(t, result.NodePath, []string{"start_1", "understanding_1", "policy_1", "policy_route_1", "retrieve_end"})
}
+func TestExecutorKnowledgeRetrieveRequiresNodeKnowledgeBases(t *testing.T) {
+ _, err := NewExecutor().Execute(context.Background(), Input{
+ Definition: knowledgeRetrieveWorkflowDefinition(nil),
+ UserMessage: models.Message{
+ Content: "产品价格",
+ },
+ AIAgent: models.AIAgent{
+ KnowledgeIDs: "1,2",
+ },
+ })
+ if err == nil {
+ t.Fatalf("expected knowledge retrieve without node knowledge bases to fail")
+ }
+ if !strings.Contains(err.Error(), "knowledge retrieve node requires knowledgeBaseIds") {
+ t.Fatalf("unexpected error: %v", err)
+ }
+}
+
func TestExecutorLLMReplyUsesAgentFallbackWhenDeclaredKnowledgeIsEmpty(t *testing.T) {
result, err := NewExecutor().Execute(context.Background(), Input{
Definition: emptyKnowledgeReplyDefinition(),
@@ -347,7 +363,6 @@ func TestExecutorLLMReplyUsesAgentFallbackWhenDeclaredKnowledgeIsEmpty(t *testin
Content: "产品功能",
},
AIAgent: models.AIAgent{
- KnowledgeIDs: "1",
FallbackMode: enums.AIAgentFallbackModeNoAnswer,
FallbackMessage: "我暂时没有找到足够准确的信息。你可以补充更具体的问题,我再继续帮你查。",
SystemPrompt: "不要编造事实。",
@@ -564,6 +579,20 @@ func conditionalReplyDefinition() dsl.Definition {
)
}
+func knowledgeRetrieveWorkflowDefinition(config any) dsl.Definition {
+ return wfTestDefinition(
+ []dsl.Node{
+ wfTestNode("start_1", workflowregistry.NodeTypeStart, "Start", nil, nil),
+ wfTestNode("retrieve_1", workflowregistry.NodeTypeKnowledgeRetrieve, "Retrieve", wfTestInputs("query", "start_1", "userMessage"), config),
+ wfTestNode("end_1", workflowregistry.NodeTypeEnd, "End", nil, nil),
+ },
+ []dsl.Edge{
+ wfTestEdge("start_1", "retrieve_1", "edge_start_retrieve"),
+ wfTestEdge("retrieve_1", "end_1", "edge_retrieve_end"),
+ },
+ )
+}
+
func ticketDraftReadyWorkflowDefinition() dsl.Definition {
return wfTestDefinition(
[]dsl.Node{
diff --git a/internal/ai/workflow/registry/registry.go b/internal/ai/workflow/registry/registry.go
index f9cb5dc..64144a3 100644
--- a/internal/ai/workflow/registry/registry.go
+++ b/internal/ai/workflow/registry/registry.go
@@ -32,7 +32,6 @@ func DefaultRegistry() *Registry {
output("messageId", "消息 ID", VariableTypeInteger, "客户本轮消息的内部编号。"),
output("aiAgentId", "AI Agent ID", VariableTypeInteger, "当前处理会话的 AI Agent 编号。"),
output("userMessage", "用户消息", VariableTypeString, "客户本轮发送的原始消息内容。"),
- output("knowledgeBaseIds", "知识库 ID 列表", VariableTypeIntegerArray, "当前 AI Agent 已绑定的知识库编号列表。"),
},
},
NodeSpec{
@@ -120,6 +119,14 @@ func DefaultRegistry() *Registry {
Description: "Retrieve knowledge for the current user message.",
Icon: "BookOpenIcon",
RiskLevel: NodeRiskLevelLow,
+ ConfigSchema: map[string]any{
+ "knowledgeBaseIds": map[string]any{
+ "type": string(VariableTypeIntegerArray),
+ "label": "知识库",
+ "required": true,
+ "description": "本节点检索时使用的知识库列表,按顺序表示优先级。",
+ },
+ },
InputSchema: []VariableSpec{
requiredInput("query", "检索问题", VariableTypeString, "用于检索知识库的客户问题或查询文本。"),
},
diff --git a/internal/ai/workflow/registry/registry_test.go b/internal/ai/workflow/registry/registry_test.go
index 3ef41ac..539067f 100644
--- a/internal/ai/workflow/registry/registry_test.go
+++ b/internal/ai/workflow/registry/registry_test.go
@@ -23,8 +23,8 @@ func TestDefaultRegistryExposesStartOutputs(t *testing.T) {
if !hasVariable(spec.OutputSchema, "userMessage", VariableTypeString) {
t.Fatalf("expected start output userMessage:string, got %#v", spec.OutputSchema)
}
- if !hasVariable(spec.OutputSchema, "knowledgeBaseIds", VariableTypeIntegerArray) {
- t.Fatalf("expected start output knowledgeBaseIds:array, got %#v", spec.OutputSchema)
+ if hasVariableName(spec.OutputSchema, "knowledgeBaseIds") {
+ t.Fatalf("did not expect start output knowledgeBaseIds after knowledge binding moved to retrieve node, got %#v", spec.OutputSchema)
}
}
@@ -39,6 +39,9 @@ func TestDefaultRegistryExposesKnowledgeRetrieveInputsAndOutputs(t *testing.T) {
if !hasVariable(spec.OutputSchema, "items", VariableTypeObjectArray) {
t.Fatalf("expected knowledge_retrieve output items:array