feat(wecom-aibot): add context management for stream tasks to improve agent cancellation

This commit is contained in:
Zhang Rui 2026-02-28 15:38:49 +08:00
parent 0b6d913dfc
commit 4e09c91dda

View file

@ -49,6 +49,8 @@ type streamTask struct {
Finished bool // fully done Finished bool // fully done
mu sync.Mutex mu sync.Mutex
answerCh chan string // receives agent reply from Send() answerCh chan string // receives agent reply from Send()
ctx context.Context // canceled when task is removed; used to interrupt the agent goroutine
cancel context.CancelFunc // call on task removal to cancel ctx
} }
// WeComAIBotMessage represents the decrypted JSON message from WeCom AI Bot // WeComAIBotMessage represents the decrypted JSON message from WeCom AI Bot
@ -237,6 +239,9 @@ func (c *WeComAIBotChannel) Send(ctx context.Context, msg bus.OutboundMessage) e
// Stream still open: deliver via answerCh for the next poll response. // Stream still open: deliver via answerCh for the next poll response.
select { select {
case task.answerCh <- msg.Content: case task.answerCh <- msg.Content:
case <-task.ctx.Done():
// Task was canceled (cleanup removed it); silently drop the reply.
return nil
case <-ctx.Done(): case <-ctx.Done():
return ctx.Err() return ctx.Err()
} }
@ -490,6 +495,10 @@ func (c *WeComAIBotChannel) handleTextMessage(
// Set a slightly shorter deadline so we can send a timeout notice before it gives up. // Set a slightly shorter deadline so we can send a timeout notice before it gives up.
deadline := time.Now().Add(30 * time.Second) deadline := time.Now().Add(30 * time.Second)
// Each task gets its own context derived from the channel lifetime context.
// Canceling taskCancel interrupts the agent goroutine when the task is removed.
taskCtx, taskCancel := context.WithCancel(c.ctx)
task := &streamTask{ task := &streamTask{
StreamID: streamID, StreamID: streamID,
ChatID: chatID, ChatID: chatID,
@ -499,6 +508,8 @@ func (c *WeComAIBotChannel) handleTextMessage(
Deadline: deadline, Deadline: deadline,
Finished: false, Finished: false,
answerCh: make(chan string, 1), answerCh: make(chan string, 1),
ctx: taskCtx,
cancel: taskCancel,
} }
c.taskMu.Lock() c.taskMu.Lock()
@ -506,8 +517,8 @@ func (c *WeComAIBotChannel) handleTextMessage(
c.chatTasks[chatID] = append(c.chatTasks[chatID], task) c.chatTasks[chatID] = append(c.chatTasks[chatID], task)
c.taskMu.Unlock() c.taskMu.Unlock()
// Publish to agent asynchronously; agent will call Send() with reply // Publish to agent asynchronously; agent will call Send() with reply.
// Use c.ctx (channel lifetime) instead of r.Context() which is canceled when the HTTP handler returns. // Use task.ctx (not c.ctx) so the agent goroutine is canceled when the task is removed.
go func() { go func() {
sender := bus.SenderInfo{ sender := bus.SenderInfo{
Platform: "wecom_aibot", Platform: "wecom_aibot",
@ -529,7 +540,7 @@ func (c *WeComAIBotChannel) handleTextMessage(
"stream_id": streamID, "stream_id": streamID,
"response_url": msg.ResponseURL, "response_url": msg.ResponseURL,
} }
c.HandleMessage(c.ctx, peer, msg.MsgID, userID, chatID, c.HandleMessage(task.ctx, peer, msg.MsgID, userID, chatID,
content, nil, metadata, sender) content, nil, metadata, sender)
}() }()
@ -800,11 +811,13 @@ func (c *WeComAIBotChannel) getStreamResponse(task *streamTask, timestamp, nonce
return c.encryptResponse(task.StreamID, timestamp, nonce, response) return c.encryptResponse(task.StreamID, timestamp, nonce, response)
} }
// removeTask removes a task from both streamTasks and chatTasks and marks it finished. // removeTask removes a task from both streamTasks and chatTasks, marks it finished,
// and cancels its context to interrupt the associated agent goroutine.
func (c *WeComAIBotChannel) removeTask(task *streamTask) { func (c *WeComAIBotChannel) removeTask(task *streamTask) {
task.mu.Lock() task.mu.Lock()
task.Finished = true task.Finished = true
task.mu.Unlock() task.mu.Unlock()
task.cancel() // interrupt agent goroutine bound to this task
c.taskMu.Lock() c.taskMu.Lock()
delete(c.streamTasks, task.StreamID) delete(c.streamTasks, task.StreamID)
@ -1114,6 +1127,7 @@ func (c *WeComAIBotChannel) cleanupOldTasks() {
for id, task := range c.streamTasks { for id, task := range c.streamTasks {
if task.CreatedTime.Before(cutoff) { if task.CreatedTime.Before(cutoff) {
delete(c.streamTasks, id) delete(c.streamTasks, id)
task.cancel() // interrupt agent goroutine still waiting for LLM
queue := c.chatTasks[task.ChatID] queue := c.chatTasks[task.ChatID]
for i, t := range queue { for i, t := range queue {
if t == task { if t == task {
@ -1130,11 +1144,14 @@ func (c *WeComAIBotChannel) cleanupOldTasks() {
} }
} }
// Also clean up StreamClosed tasks from chatTasks that are older than 1 hour. // Also clean up StreamClosed tasks from chatTasks that are older than 1 hour.
// These were removed from streamTasks earlier but kept alive for response_url delivery.
for chatID, queue := range c.chatTasks { for chatID, queue := range c.chatTasks {
filtered := queue[:0] filtered := queue[:0]
for _, t := range queue { for _, t := range queue {
if !t.Finished && t.CreatedTime.After(cutoff) { if !t.Finished && t.CreatedTime.After(cutoff) {
filtered = append(filtered, t) filtered = append(filtered, t)
} else if !t.Finished {
t.cancel() // cancel any lingering agent goroutine
} }
} }
if len(filtered) == 0 { if len(filtered) == 0 {