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{
|
||||
|
||||
Reference in New Issue
Block a user