Refactor AI Agent and AI Config handling across multiple files
- Updated function signatures to accept AI Agent and AI Config as non-pointer types for better clarity and safety. - Modified instances where AI Agent and AI Config were dereferenced to improve code readability. - Removed unnecessary nil checks for AI Agent and AI Config, simplifying the logic. - Adjusted related tests and services to align with the new function signatures. - Cleaned up code in runtime, skills, and executor packages to ensure consistency in handling AI configurations.
This commit is contained in:
@@ -22,9 +22,7 @@ func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, col
|
||||
}
|
||||
if collector != nil {
|
||||
collector.Data.Input.HistoryMessageCount = len(history.Messages)
|
||||
if req.AIAgent != nil {
|
||||
collector.Data.Input.KnowledgeBaseIDs = utils.SplitInt64s(req.AIAgent.KnowledgeIDs)
|
||||
}
|
||||
collector.Data.Input.KnowledgeBaseIDs = utils.SplitInt64s(req.AIAgent.KnowledgeIDs)
|
||||
collector.Data.Input.CurrentUserMessagePreview = preview(req.UserMessage.Content, 120)
|
||||
}
|
||||
messages := make([]*schema.Message, 0, len(history.Messages)+3)
|
||||
@@ -41,7 +39,7 @@ func buildRunMessages(ctx context.Context, req RunInput, summary *RunResult, col
|
||||
}
|
||||
|
||||
func appendRetrievedContext(ctx context.Context, req RunInput, summary *RunResult, collector *callbacks.RuntimeTraceCollector, messages *[]*schema.Message) knowledgeGuardDecision {
|
||||
if req.AIAgent == nil || req.UserMessage == nil || messages == nil {
|
||||
if req.UserMessage == nil || messages == nil {
|
||||
return knowledgeGuardDecision{}
|
||||
}
|
||||
retriever := retrievers.NewKnowledgeRetriever(req.AIAgent)
|
||||
|
||||
@@ -15,8 +15,8 @@ type knowledgeGuardDecision struct {
|
||||
Instructions []*schema.Message
|
||||
}
|
||||
|
||||
func buildKnowledgeGuardDecision(aiAgent *models.AIAgent, retrieveResult *retrievers.KnowledgeRetrieveResult) knowledgeGuardDecision {
|
||||
if aiAgent == nil || retrieveResult == nil || len(retrieveResult.KnowledgeBaseIDs) == 0 {
|
||||
func buildKnowledgeGuardDecision(aiAgent models.AIAgent, retrieveResult *retrievers.KnowledgeRetrieveResult) knowledgeGuardDecision {
|
||||
if retrieveResult == nil || len(retrieveResult.KnowledgeBaseIDs) == 0 {
|
||||
return knowledgeGuardDecision{}
|
||||
}
|
||||
fallbackReply := resolveKnowledgeFallbackReply(aiAgent, retrieveResult.FallbackMode)
|
||||
@@ -32,11 +32,9 @@ func buildKnowledgeGuardDecision(aiAgent *models.AIAgent, retrieveResult *retrie
|
||||
}
|
||||
}
|
||||
|
||||
func resolveKnowledgeFallbackReply(aiAgent *models.AIAgent, fallbackMode enums.KnowledgeFallbackMode) string {
|
||||
if aiAgent != nil {
|
||||
if reply := strings.TrimSpace(aiAgent.FallbackMessage); reply != "" {
|
||||
return reply
|
||||
}
|
||||
func resolveKnowledgeFallbackReply(aiAgent models.AIAgent, fallbackMode enums.KnowledgeFallbackMode) string {
|
||||
if reply := strings.TrimSpace(aiAgent.FallbackMessage); reply != "" {
|
||||
return reply
|
||||
}
|
||||
switch fallbackMode {
|
||||
case enums.KnowledgeFallbackModeSuggestRetry:
|
||||
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
|
||||
func TestBuildKnowledgeGuardDecisionFallsBackWhenKnowledgeMisses(t *testing.T) {
|
||||
agent := newKnowledgeGuardAgentFixture()
|
||||
decision := buildKnowledgeGuardDecision(&agent, &retrievers.KnowledgeRetrieveResult{
|
||||
decision := buildKnowledgeGuardDecision(agent, &retrievers.KnowledgeRetrieveResult{
|
||||
KnowledgeBaseIDs: []int64{1},
|
||||
FallbackMode: enums.KnowledgeFallbackModeSuggestRetry,
|
||||
})
|
||||
@@ -28,7 +28,7 @@ func TestBuildKnowledgeGuardDecisionFallsBackWhenKnowledgeMisses(t *testing.T) {
|
||||
func TestBuildKnowledgeGuardDecisionUsesAgentFallbackMessage(t *testing.T) {
|
||||
agent := newKnowledgeGuardAgentFixture()
|
||||
agent.FallbackMessage = "请联系人工客服"
|
||||
decision := buildKnowledgeGuardDecision(&agent, &retrievers.KnowledgeRetrieveResult{
|
||||
decision := buildKnowledgeGuardDecision(agent, &retrievers.KnowledgeRetrieveResult{
|
||||
KnowledgeBaseIDs: []int64{1},
|
||||
FallbackMode: enums.KnowledgeFallbackModeNoAnswer,
|
||||
})
|
||||
@@ -40,7 +40,7 @@ func TestBuildKnowledgeGuardDecisionUsesAgentFallbackMessage(t *testing.T) {
|
||||
|
||||
func TestBuildKnowledgeGuardDecisionInjectsStrictInstructionOnHit(t *testing.T) {
|
||||
agent := newKnowledgeGuardAgentFixture()
|
||||
decision := buildKnowledgeGuardDecision(&agent, &retrievers.KnowledgeRetrieveResult{
|
||||
decision := buildKnowledgeGuardDecision(agent, &retrievers.KnowledgeRetrieveResult{
|
||||
KnowledgeBaseIDs: []int64{1},
|
||||
Hits: []rag.RetrieveResult{
|
||||
{KnowledgeBaseID: 1, Score: 0.88},
|
||||
|
||||
@@ -33,7 +33,7 @@ func (s *Service) ExecuteRun(ctx context.Context, req RunInput) (*RunResult, err
|
||||
}
|
||||
collector := callbacks.NewRuntimeTraceCollector()
|
||||
collector.Data.RunID = summary.RunID
|
||||
if req.AIAgent == nil || req.Conversation == nil || req.UserMessage == nil {
|
||||
if req.Conversation == nil || req.UserMessage == nil {
|
||||
summary.Status = "error"
|
||||
summary.ErrorMessage = "invalid runtime request"
|
||||
collector.Data.Status = summary.Status
|
||||
@@ -42,15 +42,6 @@ func (s *Service) ExecuteRun(ctx context.Context, req RunInput) (*RunResult, err
|
||||
summary.TraceData = collector.Marshal()
|
||||
return summary, fmt.Errorf("%s", summary.ErrorMessage)
|
||||
}
|
||||
if req.AIConfig == nil {
|
||||
summary.Status = "error"
|
||||
summary.ErrorMessage = "ai config is nil"
|
||||
collector.Data.Status = summary.Status
|
||||
collector.Data.Error.Message = summary.ErrorMessage
|
||||
collector.Data.Error.Stage = "prepare"
|
||||
summary.TraceData = collector.Marshal()
|
||||
return summary, fmt.Errorf("%s", summary.ErrorMessage)
|
||||
}
|
||||
|
||||
toolDefs, err := factory.NewToolFactory().BuildMCPTools(req.AIAgent)
|
||||
if err != nil {
|
||||
@@ -149,24 +140,6 @@ func (s *Service) ExecuteResume(ctx context.Context, req ResumeInput) (*RunResul
|
||||
collector := callbacks.NewRuntimeTraceCollector()
|
||||
collector.Data.RunID = summary.RunID
|
||||
collector.Data.Interrupt.CheckPointID = summary.CheckPointID
|
||||
if req.AIAgent == nil {
|
||||
summary.Status = "error"
|
||||
summary.ErrorMessage = "ai agent is nil"
|
||||
collector.Data.Status = summary.Status
|
||||
collector.Data.Error.Message = summary.ErrorMessage
|
||||
collector.Data.Error.Stage = "resume_prepare"
|
||||
summary.TraceData = collector.Marshal()
|
||||
return summary, fmt.Errorf("%s", summary.ErrorMessage)
|
||||
}
|
||||
if req.AIConfig == nil {
|
||||
summary.Status = "error"
|
||||
summary.ErrorMessage = "ai config is nil"
|
||||
collector.Data.Status = summary.Status
|
||||
collector.Data.Error.Message = summary.ErrorMessage
|
||||
collector.Data.Error.Stage = "resume_prepare"
|
||||
summary.TraceData = collector.Marshal()
|
||||
return summary, fmt.Errorf("%s", summary.ErrorMessage)
|
||||
}
|
||||
if summary.CheckPointID == "" {
|
||||
summary.Status = "error"
|
||||
summary.ErrorMessage = "checkpoint id is required"
|
||||
|
||||
@@ -8,8 +8,8 @@ import (
|
||||
type RunInput struct {
|
||||
Conversation *models.Conversation
|
||||
UserMessage *models.Message
|
||||
AIAgent *models.AIAgent
|
||||
AIConfig *models.AIConfig
|
||||
AIAgent models.AIAgent
|
||||
AIConfig models.AIConfig
|
||||
SelectedSkill *models.SkillDefinition
|
||||
SkillRouteReason string
|
||||
SkillRouteTrace string
|
||||
@@ -19,8 +19,8 @@ type RunInput struct {
|
||||
|
||||
type ResumeInput struct {
|
||||
Conversation *models.Conversation
|
||||
AIAgent *models.AIAgent
|
||||
AIConfig *models.AIConfig
|
||||
AIAgent models.AIAgent
|
||||
AIConfig models.AIConfig
|
||||
CheckPointID string
|
||||
ResumeData map[string]string
|
||||
ToolSet *registry.ToolSet
|
||||
|
||||
Reference in New Issue
Block a user