refactor(runtime): update function signatures to use value receivers for models

This commit is contained in:
mlogclub
2026-04-19 11:38:05 +08:00
parent 680793f527
commit 3e98e9379c
9 changed files with 107 additions and 95 deletions
+37 -41
View File
@@ -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
}