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:
mlogclub
2026-04-17 17:57:01 +08:00
parent 3b062c327c
commit 976b9defde
36 changed files with 102 additions and 282 deletions
@@ -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},
+1 -28
View File
@@ -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"
+4 -4
View File
@@ -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