feat(ops): 优雅停机 drain —— 三 Go 服务 SIGTERM 后排空在途,不硬切
滚动更新/重启时旧的"信号到→进程退"会硬切在途工作:dispatcher 在途任务被掐断、 mcp-go 在途工具调用让 dispatcher 干等超时、gateway 在途 HTTP 请求被截断。 共享 bus 加在途追踪 + drain: - ConsumeTasks/ServeTool 各加 sync.WaitGroup 跟踪在途 goroutine,返回 drain(ctx): 先停止接新活(cc.Stop / Unsubscribe),再等在途跑完至 drain 超时。 - 关键修复:任务 handler 的 ctx 改为派生自 context.Background()(而非信号 ctx), 否则 SIGTERM 会立即取消在途任务的 ctx,drain 形同虚设。超时未跑完才由 JetStream AckWait 重投兜底(不丢任务)。 - DrainTimeout():SHUTDOWN_DRAIN_TIMEOUT 秒,默认 30s。 各服务收尾: - gateway:r.Run → http.Server + signal.NotifyContext + srv.Shutdown(drain 在途请求), 随后 defer 关 db/redis/bus(HTTP 排空后才断后端连接)。 - dispatcher:收到信号 → drain 在途任务跑完再退。 - mcp-go:收到信号 → drain 在途工具调用回完再退(dispatcher 拿到结果而非超时)。 bus 加 TestGracefulDrain(drain 等满在途任务 + 验证在途 ctx 不被取消),e2e 测试适配 新签名,四模块全绿。live:在途任务执行中 kill -TERM dispatcher → 日志「drain 在途任务」 → 11s 后 task done(511字完整生成) → 「drain 完成,退出」,任务终态 done 非 failed; gateway/mcp-go 同样优雅退出。 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -9,6 +9,7 @@ import (
|
||||
"log"
|
||||
"os"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/nats-io/nats.go"
|
||||
@@ -246,14 +247,17 @@ func toolConcurrency() int {
|
||||
// ServeTool 以队列组订阅工具主题(可用通配 sundynix.tools.go.>),队列组内多副本自动负载均衡(水平扩)。
|
||||
// 单实例内每个请求分发到独立 goroutine 并发处理(垂直扩,上限 MCP_TOOL_CONCURRENCY)——
|
||||
// 回调立即返回不阻塞 NATS 投递;信号量约束同时执行的 handler 数。返回的 unsub 用于退订。
|
||||
func (b *Bus) ServeTool(subject, queue string, h ToolHandler) (unsub func() error, err error) {
|
||||
func (b *Bus) ServeTool(subject, queue string, h ToolHandler) (drain func(context.Context), err error) {
|
||||
sem := make(chan struct{}, toolConcurrency())
|
||||
var wg sync.WaitGroup // 跟踪在途工具调用,供优雅停机 drain 等待回完
|
||||
sub, err := b.nc.QueueSubscribe(subject, queue, func(m *nats.Msg) {
|
||||
// 立即起 goroutine 并返回,让 NATS 继续投递下一条(核心 NATS 单订阅回调本是串行)。
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
sem <- struct{}{} // 限并发:超额则在此短暂排队
|
||||
defer func() {
|
||||
<-sem
|
||||
wg.Done()
|
||||
if r := recover(); r != nil { // handler panic → 回错误结果,调用方不必干等超时
|
||||
respond(m, &contract.ToolResult{OK: false, Error: fmt.Sprintf("tool panic: %v", r)})
|
||||
}
|
||||
@@ -279,7 +283,11 @@ func (b *Bus) ServeTool(subject, queue string, h ToolHandler) (unsub func() erro
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("serve tool %s: %w", subject, err)
|
||||
}
|
||||
return sub.Unsubscribe, nil
|
||||
// drain:退订(停止接新请求),等在途工具调用回完(至多 dctx)。
|
||||
return func(dctx context.Context) {
|
||||
_ = sub.Unsubscribe()
|
||||
drainWait(&wg, dctx)
|
||||
}, nil
|
||||
}
|
||||
|
||||
func respond(m *nats.Msg, res *contract.ToolResult) {
|
||||
@@ -523,10 +531,31 @@ func taskConcurrency() int {
|
||||
return 8
|
||||
}
|
||||
|
||||
// DrainTimeout 是优雅停机时等待在途工作(任务/工具调用/HTTP 请求)跑完的上限。
|
||||
// 经 SHUTDOWN_DRAIN_TIMEOUT 秒配置,缺省 30s。超时即放弃等待退出(在途未 ack 由 JetStream 重投兜底)。
|
||||
func DrainTimeout() time.Duration {
|
||||
if v := os.Getenv("SHUTDOWN_DRAIN_TIMEOUT"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil && n > 0 {
|
||||
return time.Duration(n) * time.Second
|
||||
}
|
||||
}
|
||||
return 30 * time.Second
|
||||
}
|
||||
|
||||
// drainWait 停止接新活后,等待在途 WaitGroup 跑完,或 dctx 到期放弃。
|
||||
func drainWait(wg *sync.WaitGroup, dctx context.Context) {
|
||||
done := make(chan struct{})
|
||||
go func() { wg.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-dctx.Done():
|
||||
}
|
||||
}
|
||||
|
||||
// ConsumeTasks 在持久消费者上消费任务,队列组内负载均衡。
|
||||
// 每个任务分发到独立 worker goroutine 并发执行——一个慢任务/HITL 待审不再阻塞后续任务。
|
||||
// 并发上限由信号量 + 消费者 MaxAckPending 双重约束(背压)。返回的 stop 用于优雅停止消费。
|
||||
func (b *Bus) ConsumeTasks(ctx context.Context, h TaskHandler) (stop func(), err error) {
|
||||
func (b *Bus) ConsumeTasks(ctx context.Context, h TaskHandler) (drain func(context.Context), err error) {
|
||||
concurrency := taskConcurrency()
|
||||
cons, err := b.js.CreateOrUpdateConsumer(ctx, contract.StreamTasks, jetstream.ConsumerConfig{
|
||||
Durable: contract.ConsumerDurable,
|
||||
@@ -542,6 +571,7 @@ func (b *Bus) ConsumeTasks(ctx context.Context, h TaskHandler) (stop func(), err
|
||||
return nil, fmt.Errorf("create consumer: %w", err)
|
||||
}
|
||||
sem := make(chan struct{}, concurrency) // 限并发:最多 N 个任务同时执行
|
||||
var wg sync.WaitGroup // 跟踪在途任务,供优雅停机 drain 等待跑完
|
||||
cc, err := cons.Consume(func(msg jetstream.Msg) {
|
||||
t, err := contract.Unmarshal(msg.Data())
|
||||
if err != nil {
|
||||
@@ -554,17 +584,20 @@ func (b *Bus) ConsumeTasks(ctx context.Context, h TaskHandler) (stop func(), err
|
||||
case <-ctx.Done():
|
||||
return
|
||||
}
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer func() {
|
||||
<-sem // 释放并发额度
|
||||
wg.Done()
|
||||
if r := recover(); r != nil {
|
||||
// 任务处理 panic:丢弃不重投(避免崩溃循环),记录后继续。
|
||||
log.Printf("[bus] task %s handler panic: %v", t.ID, r)
|
||||
_ = msg.Term()
|
||||
}
|
||||
}()
|
||||
// 从消息头还原上游链路,开消费 span(成为 gateway 发布 span 的子节点)。
|
||||
mctx := extractTrace(ctx, nats.Header(msg.Headers()))
|
||||
// 关键:handler ctx 派生自 Background(而非信号 ctx),使「停止消费」不会立刻掐断在途任务——
|
||||
// 在途任务由 drain 等待至完成;drain 超时未跑完才由 JetStream AckWait 重投兜底。
|
||||
mctx := extractTrace(context.Background(), nats.Header(msg.Headers()))
|
||||
mctx, span := tracer().Start(mctx, "nats.consume task",
|
||||
trace.WithSpanKind(trace.SpanKindConsumer),
|
||||
trace.WithAttributes(attribute.String("sundynix.task_id", t.ID)))
|
||||
@@ -580,5 +613,9 @@ func (b *Bus) ConsumeTasks(ctx context.Context, h TaskHandler) (stop func(), err
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("consume: %w", err)
|
||||
}
|
||||
return cc.Stop, nil
|
||||
// drain:停止新投递,等在途任务跑完(至多 dctx)。
|
||||
return func(dctx context.Context) {
|
||||
cc.Stop()
|
||||
drainWait(&wg, dctx)
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -55,14 +55,14 @@ func TestTaskRoundTrip(t *testing.T) {
|
||||
defer dp.Close()
|
||||
|
||||
got := make(chan *contract.Task, 1)
|
||||
stop, err := dp.ConsumeTasks(ctx, func(_ context.Context, task *contract.Task) error {
|
||||
drain, err := dp.ConsumeTasks(ctx, func(_ context.Context, task *contract.Task) error {
|
||||
got <- task
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("consume: %v", err)
|
||||
}
|
||||
defer stop()
|
||||
defer drain(context.Background())
|
||||
|
||||
// --- Gateway 发布一个任务 ---
|
||||
want := &contract.Task{
|
||||
@@ -104,7 +104,7 @@ func TestToolCallRoundTrip(t *testing.T) {
|
||||
}
|
||||
defer srv.Close()
|
||||
|
||||
unsub, err := srv.ServeTool(contract.SubjectToolsGoAll, contract.QueueToolsGo,
|
||||
drain, err := srv.ServeTool(contract.SubjectToolsGoAll, contract.QueueToolsGo,
|
||||
func(_ context.Context, call *contract.ToolCall) *contract.ToolResult {
|
||||
if call.Tool != "wiki_search" {
|
||||
return &contract.ToolResult{OK: false, Error: "unknown tool"}
|
||||
@@ -114,7 +114,7 @@ func TestToolCallRoundTrip(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("serve tool: %v", err)
|
||||
}
|
||||
defer func() { _ = unsub() }()
|
||||
defer drain(context.Background())
|
||||
|
||||
// --- Dispatcher 侧:同步调用工具 ---
|
||||
dp, err := bus.Connect(url)
|
||||
@@ -220,7 +220,7 @@ func TestConcurrentConsume(t *testing.T) {
|
||||
|
||||
blockA := make(chan struct{})
|
||||
doneB := make(chan string, 1)
|
||||
stop, err := dp.ConsumeTasks(ctx, func(_ context.Context, task *contract.Task) error {
|
||||
drain, err := dp.ConsumeTasks(ctx, func(_ context.Context, task *contract.Task) error {
|
||||
switch task.ID {
|
||||
case "A":
|
||||
<-blockA // 模拟审批长阻塞
|
||||
@@ -232,7 +232,7 @@ func TestConcurrentConsume(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("consume: %v", err)
|
||||
}
|
||||
defer stop()
|
||||
defer drain(context.Background())
|
||||
defer close(blockA)
|
||||
|
||||
// 先发 A 并等它被消费、卡在 handler 里;再发 B。
|
||||
@@ -267,7 +267,7 @@ func TestConcurrentToolServe(t *testing.T) {
|
||||
defer srv.Close()
|
||||
|
||||
blockSlow := make(chan struct{})
|
||||
unsub, err := srv.ServeTool(contract.SubjectToolsGoAll, contract.QueueToolsGo,
|
||||
drain, err := srv.ServeTool(contract.SubjectToolsGoAll, contract.QueueToolsGo,
|
||||
func(_ context.Context, call *contract.ToolCall) *contract.ToolResult {
|
||||
if call.Tool == "slow" {
|
||||
<-blockSlow // 模拟慢工具(如 RAG 嵌入)
|
||||
@@ -277,7 +277,7 @@ func TestConcurrentToolServe(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("serve tool: %v", err)
|
||||
}
|
||||
defer unsub()
|
||||
defer drain(context.Background())
|
||||
defer close(blockSlow)
|
||||
|
||||
cli, err := bus.Connect(url)
|
||||
@@ -306,3 +306,57 @@ func TestConcurrentToolServe(t *testing.T) {
|
||||
}
|
||||
t.Log("✓ 单实例工具并发生效:slow 仍阻塞时 fast 已返回")
|
||||
}
|
||||
|
||||
// TestGracefulDrain 验证优雅停机:drain 会等在途任务跑完(不被信号掐断),
|
||||
// 且在途任务的 ctx 不随消费停止而取消。
|
||||
func TestGracefulDrain(t *testing.T) {
|
||||
url := startEmbeddedNATS(t)
|
||||
gw, err := bus.Connect(url)
|
||||
if err != nil {
|
||||
t.Fatalf("gateway connect: %v", err)
|
||||
}
|
||||
defer gw.Close()
|
||||
if err := gw.EnsureTaskStream(context.Background()); err != nil {
|
||||
t.Fatalf("ensure stream: %v", err)
|
||||
}
|
||||
dp, err := bus.Connect(url)
|
||||
if err != nil {
|
||||
t.Fatalf("dispatcher connect: %v", err)
|
||||
}
|
||||
defer dp.Close()
|
||||
|
||||
started := make(chan struct{})
|
||||
finished := make(chan struct{})
|
||||
drain, err := dp.ConsumeTasks(context.Background(), func(c context.Context, _ *contract.Task) error {
|
||||
close(started)
|
||||
time.Sleep(500 * time.Millisecond) // 模拟在途长任务
|
||||
if c.Err() != nil { // 关键:消费停止不应取消在途任务的 ctx
|
||||
t.Errorf("在途任务 ctx 不应被取消: %v", c.Err())
|
||||
}
|
||||
close(finished)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("consume: %v", err)
|
||||
}
|
||||
if _, err := gw.PublishTask(context.Background(), &contract.Task{ID: "drain1", Graph: json.RawMessage(`{}`)}); err != nil {
|
||||
t.Fatalf("publish: %v", err)
|
||||
}
|
||||
<-started // 等任务进入 handler
|
||||
|
||||
t0 := time.Now()
|
||||
dctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
drain(dctx) // 应阻塞至在途任务跑完
|
||||
elapsed := time.Since(t0)
|
||||
|
||||
select {
|
||||
case <-finished:
|
||||
default:
|
||||
t.Fatal("drain 返回时在途任务竟未完成(被掐断或未等待)")
|
||||
}
|
||||
if elapsed < 300*time.Millisecond {
|
||||
t.Fatalf("drain 过早返回(%v),未真正等待在途任务", elapsed)
|
||||
}
|
||||
t.Logf("✓ drain 等待在途任务跑完:耗时 %v", elapsed)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user