refactor(runtime): update function signatures to use value receivers for models
This commit is contained in:
@@ -0,0 +1,21 @@
|
||||
package runtime
|
||||
|
||||
import (
|
||||
applicationruntime "cs-agent/internal/ai/application/runtime"
|
||||
"cs-agent/internal/models"
|
||||
)
|
||||
|
||||
type aiReplyContext struct {
|
||||
Conversation models.Conversation
|
||||
Message models.Message
|
||||
AIAgent models.AIAgent
|
||||
Trace *aiReplyTraceData
|
||||
SummaryRef **applicationruntime.Summary
|
||||
PendingInterrupt *models.ConversationInterrupt
|
||||
}
|
||||
|
||||
func (c aiReplyContext) setSummary(summary *applicationruntime.Summary) {
|
||||
if c.SummaryRef != nil {
|
||||
*c.SummaryRef = summary
|
||||
}
|
||||
}
|
||||
@@ -6,7 +6,6 @@ import (
|
||||
|
||||
applicationruntime "cs-agent/internal/ai/application/runtime"
|
||||
"cs-agent/internal/ai/runtime/graphs"
|
||||
"cs-agent/internal/models"
|
||||
svc "cs-agent/internal/services"
|
||||
)
|
||||
|
||||
@@ -16,31 +15,30 @@ func newReplyInterruptService() *replyInterruptService {
|
||||
return &replyInterruptService{}
|
||||
}
|
||||
|
||||
func (s *replyInterruptService) ResumePendingInterrupt(ctx context.Context, owner *aiReplyService, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
|
||||
pendingInterrupt *models.ConversationInterrupt, trace *aiReplyTraceData, summaryRef **applicationruntime.Summary) error {
|
||||
if pendingInterrupt == nil || owner == nil {
|
||||
func (s *replyInterruptService) ResumePendingInterrupt(ctx context.Context, owner *aiReplyService, replyCtx aiReplyContext) error {
|
||||
if replyCtx.PendingInterrupt == nil || owner == nil {
|
||||
return nil
|
||||
}
|
||||
summary, err := owner.executor.ResumePendingInterrupt(ctx, runtimeReplyResumeInput{
|
||||
Conversation: conversation,
|
||||
Message: message,
|
||||
AIAgent: aiAgent,
|
||||
PendingInterrupt: pendingInterrupt,
|
||||
Trace: trace,
|
||||
Conversation: replyCtx.Conversation,
|
||||
Message: replyCtx.Message,
|
||||
AIAgent: replyCtx.AIAgent,
|
||||
PendingInterrupt: replyCtx.PendingInterrupt,
|
||||
Trace: replyCtx.Trace,
|
||||
})
|
||||
*summaryRef = summary
|
||||
replyCtx.setSummary(summary)
|
||||
if err != nil {
|
||||
if isCheckpointMissingError(err) {
|
||||
summary = expiredInterruptSummary()
|
||||
*summaryRef = summary
|
||||
trace.Status = "interrupt_expired"
|
||||
trace.FinalAction = "expired"
|
||||
replyCtx.setSummary(summary)
|
||||
replyCtx.Trace.Status = "interrupt_expired"
|
||||
replyCtx.Trace.FinalAction = "expired"
|
||||
replyMessage, expireErr := owner.commit.CommitAIReply(replyCommitInput{
|
||||
Conversation: conversation,
|
||||
Message: message,
|
||||
AIAgent: aiAgent,
|
||||
Conversation: replyCtx.Conversation,
|
||||
Message: replyCtx.Message,
|
||||
AIAgent: replyCtx.AIAgent,
|
||||
ReplyText: summary.ReplyText,
|
||||
Trace: trace,
|
||||
Trace: replyCtx.Trace,
|
||||
ClientPrefix: "ai_interrupt_expired",
|
||||
})
|
||||
if expireErr != nil {
|
||||
@@ -50,7 +48,7 @@ func (s *replyInterruptService) ResumePendingInterrupt(ctx context.Context, owne
|
||||
if replyMessage != nil {
|
||||
lastResumeMessageID = replyMessage.ID
|
||||
}
|
||||
if expireMarkErr := svc.ConversationInterruptService.MarkExpired(pendingInterrupt.ID, lastResumeMessageID); expireMarkErr != nil {
|
||||
if expireMarkErr := svc.ConversationInterruptService.MarkExpired(replyCtx.PendingInterrupt.ID, lastResumeMessageID); expireMarkErr != nil {
|
||||
return expireMarkErr
|
||||
}
|
||||
return nil
|
||||
@@ -58,15 +56,15 @@ func (s *replyInterruptService) ResumePendingInterrupt(ctx context.Context, owne
|
||||
return err
|
||||
}
|
||||
if summary != nil && summary.Interrupted {
|
||||
return s.HandleInterruptedResume(owner, conversation, message, aiAgent, pendingInterrupt, summary, trace)
|
||||
return s.HandleInterruptedResume(owner, replyCtx, summary)
|
||||
}
|
||||
if summary != nil && strings.TrimSpace(summary.ReplyText) != "" {
|
||||
replyMessage, err := owner.commit.CommitAIReply(replyCommitInput{
|
||||
Conversation: conversation,
|
||||
Message: message,
|
||||
AIAgent: aiAgent,
|
||||
Conversation: replyCtx.Conversation,
|
||||
Message: replyCtx.Message,
|
||||
AIAgent: replyCtx.AIAgent,
|
||||
ReplyText: summary.ReplyText,
|
||||
Trace: trace,
|
||||
Trace: replyCtx.Trace,
|
||||
ClientPrefix: "ai_resume",
|
||||
})
|
||||
if err != nil {
|
||||
@@ -77,30 +75,29 @@ func (s *replyInterruptService) ResumePendingInterrupt(ctx context.Context, owne
|
||||
replyMessageID = replyMessage.ID
|
||||
}
|
||||
if graphs.IsCancellationReply(summary.ReplyText) {
|
||||
return svc.ConversationInterruptService.MarkCancelled(pendingInterrupt.ID, replyMessageID)
|
||||
return svc.ConversationInterruptService.MarkCancelled(replyCtx.PendingInterrupt.ID, replyMessageID)
|
||||
}
|
||||
return svc.ConversationInterruptService.MarkResolved(pendingInterrupt.ID, replyMessageID)
|
||||
return svc.ConversationInterruptService.MarkResolved(replyCtx.PendingInterrupt.ID, replyMessageID)
|
||||
}
|
||||
return svc.ConversationInterruptService.MarkResolved(pendingInterrupt.ID, 0)
|
||||
return svc.ConversationInterruptService.MarkResolved(replyCtx.PendingInterrupt.ID, 0)
|
||||
}
|
||||
|
||||
func (s *replyInterruptService) HandleInterruptedSummary(owner *aiReplyService, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
|
||||
summary *applicationruntime.Summary, trace *aiReplyTraceData) error {
|
||||
func (s *replyInterruptService) HandleInterruptedSummary(owner *aiReplyService, replyCtx aiReplyContext, summary *applicationruntime.Summary) error {
|
||||
if owner == nil {
|
||||
return nil
|
||||
}
|
||||
pending := buildConversationInterrupt(conversation, message, aiAgent, summary)
|
||||
pending := buildConversationInterrupt(replyCtx.Conversation, replyCtx.Message, replyCtx.AIAgent, summary)
|
||||
if err := svc.ConversationInterruptService.CreateOrUpdatePending(pending); err != nil {
|
||||
return err
|
||||
}
|
||||
pending = svc.ConversationInterruptService.GetByCheckPointID(summary.CheckPointID)
|
||||
replyText := resolveInterruptPrompt(summary)
|
||||
replyMessage, err := owner.commit.CommitAIReply(replyCommitInput{
|
||||
Conversation: conversation,
|
||||
Message: message,
|
||||
AIAgent: aiAgent,
|
||||
Conversation: replyCtx.Conversation,
|
||||
Message: replyCtx.Message,
|
||||
AIAgent: replyCtx.AIAgent,
|
||||
ReplyText: replyText,
|
||||
Trace: trace,
|
||||
Trace: replyCtx.Trace,
|
||||
ClientPrefix: "ai_interrupt",
|
||||
})
|
||||
if err != nil {
|
||||
@@ -112,25 +109,24 @@ func (s *replyInterruptService) HandleInterruptedSummary(owner *aiReplyService,
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *replyInterruptService) HandleInterruptedResume(owner *aiReplyService, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
|
||||
pendingInterrupt *models.ConversationInterrupt, summary *applicationruntime.Summary, trace *aiReplyTraceData) error {
|
||||
if pendingInterrupt == nil || owner == nil {
|
||||
func (s *replyInterruptService) HandleInterruptedResume(owner *aiReplyService, replyCtx aiReplyContext, summary *applicationruntime.Summary) error {
|
||||
if replyCtx.PendingInterrupt == nil || owner == nil {
|
||||
return nil
|
||||
}
|
||||
replyText := resolveInterruptPrompt(summary)
|
||||
replyMessage, err := owner.commit.CommitAIReply(replyCommitInput{
|
||||
Conversation: conversation,
|
||||
Message: message,
|
||||
AIAgent: aiAgent,
|
||||
Conversation: replyCtx.Conversation,
|
||||
Message: replyCtx.Message,
|
||||
AIAgent: replyCtx.AIAgent,
|
||||
ReplyText: replyText,
|
||||
Trace: trace,
|
||||
Trace: replyCtx.Trace,
|
||||
ClientPrefix: "ai_interrupt_resume",
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if replyMessage != nil {
|
||||
return svc.ConversationInterruptService.MarkPendingAgain(pendingInterrupt.ID, firstInterruptID(summary), replyText, replyMessage.ID)
|
||||
return svc.ConversationInterruptService.MarkPendingAgain(replyCtx.PendingInterrupt.ID, firstInterruptID(summary), replyText, replyMessage.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -46,6 +46,13 @@ func (s *aiReplyService) TriggerReply(ctx context.Context, conversation models.C
|
||||
startedAt := time.Now()
|
||||
trace := &aiReplyTraceData{Status: "started"}
|
||||
var summary *applicationruntime.Summary
|
||||
replyCtx := aiReplyContext{
|
||||
Conversation: conversation,
|
||||
Message: message,
|
||||
AIAgent: aiAgent,
|
||||
Trace: trace,
|
||||
SummaryRef: &summary,
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -65,46 +72,43 @@ func (s *aiReplyService) TriggerReply(ctx context.Context, conversation models.C
|
||||
})
|
||||
}()
|
||||
if pendingInterrupt := svc.ConversationInterruptService.FindLatestPendingByConversationID(conversation.ID); pendingInterrupt != nil {
|
||||
return s.resumePendingInterrupt(ctx, conversation, message, aiAgent, pendingInterrupt, trace, &summary)
|
||||
replyCtx.PendingInterrupt = pendingInterrupt
|
||||
return s.resumePendingInterrupt(ctx, replyCtx)
|
||||
}
|
||||
return s.executeReply(ctx, conversation, message, aiAgent, trace, &summary)
|
||||
return s.executeReply(ctx, replyCtx)
|
||||
}
|
||||
|
||||
func (s *aiReplyService) resumePendingInterrupt(ctx context.Context, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
|
||||
pendingInterrupt *models.ConversationInterrupt, trace *aiReplyTraceData, summaryRef **applicationruntime.Summary) error {
|
||||
return s.interrupts.ResumePendingInterrupt(ctx, s, conversation, message, aiAgent, pendingInterrupt, trace, summaryRef)
|
||||
func (s *aiReplyService) resumePendingInterrupt(ctx context.Context, replyCtx aiReplyContext) error {
|
||||
return s.interrupts.ResumePendingInterrupt(ctx, s, replyCtx)
|
||||
}
|
||||
|
||||
func (s *aiReplyService) executeReply(ctx context.Context, conversation models.Conversation, message models.Message, aiAgent models.AIAgent,
|
||||
trace *aiReplyTraceData, summaryRef **applicationruntime.Summary) error {
|
||||
func (s *aiReplyService) executeReply(ctx context.Context, replyCtx aiReplyContext) error {
|
||||
summary, err := s.executor.Run(ctx, runtimeReplyRunInput{
|
||||
Conversation: conversation,
|
||||
Message: message,
|
||||
AIAgent: aiAgent,
|
||||
Trace: trace,
|
||||
Conversation: replyCtx.Conversation,
|
||||
Message: replyCtx.Message,
|
||||
AIAgent: replyCtx.AIAgent,
|
||||
Trace: replyCtx.Trace,
|
||||
})
|
||||
if summaryRef != nil {
|
||||
*summaryRef = summary
|
||||
}
|
||||
replyCtx.setSummary(summary)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if summary != nil && summary.Interrupted {
|
||||
return s.interrupts.HandleInterruptedSummary(s, conversation, message, aiAgent, summary, trace)
|
||||
return s.interrupts.HandleInterruptedSummary(s, replyCtx, summary)
|
||||
}
|
||||
if summary != nil && strings.TrimSpace(summary.ReplyText) != "" {
|
||||
replyMessage, err := s.commit.CommitAIReply(replyCommitInput{
|
||||
Conversation: conversation,
|
||||
Message: message,
|
||||
AIAgent: aiAgent,
|
||||
Conversation: replyCtx.Conversation,
|
||||
Message: replyCtx.Message,
|
||||
AIAgent: replyCtx.AIAgent,
|
||||
ReplyText: summary.ReplyText,
|
||||
Trace: trace,
|
||||
Trace: replyCtx.Trace,
|
||||
ClientPrefix: "ai_reply",
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
trace.ReplySent = replyMessage != nil
|
||||
replyCtx.Trace.ReplySent = replyMessage != nil
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user