05c25d7099
滚动更新/重启时旧的"信号到→进程退"会硬切在途工作: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>
363 lines
10 KiB
Go
363 lines
10 KiB
Go
package bus_test
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"testing"
|
||
"time"
|
||
|
||
natsserver "github.com/nats-io/nats-server/v2/server"
|
||
natstest "github.com/nats-io/nats-server/v2/test"
|
||
|
||
"github.com/sundynix/sundynix-shared/bus"
|
||
"github.com/sundynix/sundynix-shared/contract"
|
||
)
|
||
|
||
// startEmbeddedNATS 启动一个内嵌、开启 JetStream 的 NATS 服务器,免 Docker。
|
||
func startEmbeddedNATS(t *testing.T) string {
|
||
t.Helper()
|
||
opts := natstest.DefaultTestOptions
|
||
opts.Port = -1 // 随机端口
|
||
opts.JetStream = true
|
||
opts.StoreDir = t.TempDir()
|
||
srv := natstest.RunServer(&opts)
|
||
if !srv.ReadyForConnections(5 * time.Second) {
|
||
t.Fatal("embedded NATS not ready")
|
||
}
|
||
t.Cleanup(srv.Shutdown)
|
||
_ = natsserver.Server{} // 触发包引用
|
||
return srv.ClientURL()
|
||
}
|
||
|
||
// TestTaskRoundTrip 模拟 Gateway 发布 → NATS → Dispatcher 消费 的完整任务流。
|
||
func TestTaskRoundTrip(t *testing.T) {
|
||
url := startEmbeddedNATS(t)
|
||
|
||
// --- Gateway 侧:连接并声明任务流 ---
|
||
gw, err := bus.Connect(url)
|
||
if err != nil {
|
||
t.Fatalf("gateway connect: %v", err)
|
||
}
|
||
defer gw.Close()
|
||
|
||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||
defer cancel()
|
||
|
||
if err := gw.EnsureTaskStream(ctx); err != nil {
|
||
t.Fatalf("ensure stream: %v", err)
|
||
}
|
||
|
||
// --- Dispatcher 侧:连接并开始消费 ---
|
||
dp, err := bus.Connect(url)
|
||
if err != nil {
|
||
t.Fatalf("dispatcher connect: %v", err)
|
||
}
|
||
defer dp.Close()
|
||
|
||
got := make(chan *contract.Task, 1)
|
||
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 drain(context.Background())
|
||
|
||
// --- Gateway 发布一个任务 ---
|
||
want := &contract.Task{
|
||
ID: "task_demo_001",
|
||
Graph: json.RawMessage(`{"nodes":[{"id":"n1","type":"agent"}],"edges":[]}`),
|
||
Meta: map[string]any{"user": "wt"},
|
||
}
|
||
seq, err := gw.PublishTask(ctx, want)
|
||
if err != nil {
|
||
t.Fatalf("publish: %v", err)
|
||
}
|
||
if seq == 0 {
|
||
t.Fatal("expected non-zero stream sequence")
|
||
}
|
||
|
||
// --- 断言 Dispatcher 收到同一个任务 ---
|
||
select {
|
||
case task := <-got:
|
||
if task.ID != want.ID {
|
||
t.Fatalf("task id = %q, want %q", task.ID, want.ID)
|
||
}
|
||
if task.Meta["user"] != "wt" {
|
||
t.Fatalf("task meta lost: %+v", task.Meta)
|
||
}
|
||
t.Logf("✓ 任务流打通:Gateway publish (seq=%d) → NATS → Dispatcher consume,task_id=%s", seq, task.ID)
|
||
case <-time.After(5 * time.Second):
|
||
t.Fatal("timeout: dispatcher 未收到任务")
|
||
}
|
||
}
|
||
|
||
// TestToolCallRoundTrip 模拟 Dispatcher 经 NATS 调用 → mcp-go 响应 的工具调用闭环。
|
||
func TestToolCallRoundTrip(t *testing.T) {
|
||
url := startEmbeddedNATS(t)
|
||
|
||
// --- mcp-go 侧:以队列组订阅工具主题并响应 ---
|
||
srv, err := bus.Connect(url)
|
||
if err != nil {
|
||
t.Fatalf("mcp connect: %v", err)
|
||
}
|
||
defer srv.Close()
|
||
|
||
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"}
|
||
}
|
||
return &contract.ToolResult{OK: true, Content: "命中:" + call.Args["q"].(string)}
|
||
})
|
||
if err != nil {
|
||
t.Fatalf("serve tool: %v", err)
|
||
}
|
||
defer drain(context.Background())
|
||
|
||
// --- Dispatcher 侧:同步调用工具 ---
|
||
dp, err := bus.Connect(url)
|
||
if err != nil {
|
||
t.Fatalf("dispatcher connect: %v", err)
|
||
}
|
||
defer dp.Close()
|
||
|
||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||
defer cancel()
|
||
|
||
res, err := dp.CallTool(ctx, contract.ToolSubjectGo("wiki_search"), &contract.ToolCall{
|
||
Tool: "wiki_search",
|
||
TaskID: "task_tool_001",
|
||
Args: map[string]any{"q": "向量检索"},
|
||
})
|
||
if err != nil {
|
||
t.Fatalf("call tool: %v", err)
|
||
}
|
||
if !res.OK || res.Content != "命中:向量检索" {
|
||
t.Fatalf("tool result = %+v, want ok content=命中:向量检索", res)
|
||
}
|
||
t.Logf("✓ 工具调用闭环:Dispatcher → sundynix.tools.go.wiki_search → mcp-go → %q", res.Content)
|
||
}
|
||
|
||
// TestTokenStreamRoundTrip 模拟 Dispatcher 回流 Token → Gateway 订阅 的流式闭环。
|
||
func TestTokenStreamRoundTrip(t *testing.T) {
|
||
url := startEmbeddedNATS(t)
|
||
|
||
// Gateway 侧:先订阅(core NATS 无持久化,须先连)。
|
||
gw, err := bus.Connect(url)
|
||
if err != nil {
|
||
t.Fatalf("gateway connect: %v", err)
|
||
}
|
||
defer gw.Close()
|
||
|
||
const taskID = "task_stream_001"
|
||
var got []string
|
||
done := make(chan struct{})
|
||
unsub, err := gw.SubscribeTokens(taskID,
|
||
func(tok []byte) { got = append(got, string(tok)) },
|
||
func() { close(done) },
|
||
)
|
||
if err != nil {
|
||
t.Fatalf("subscribe tokens: %v", err)
|
||
}
|
||
defer func() { _ = unsub() }()
|
||
|
||
// Dispatcher 侧:逐 Token 回流后发结束信号。
|
||
dp, err := bus.Connect(url)
|
||
if err != nil {
|
||
t.Fatalf("dispatcher connect: %v", err)
|
||
}
|
||
defer dp.Close()
|
||
|
||
want := []string{"Hello", " ", "Agent", "!"}
|
||
for _, tok := range want {
|
||
if err := dp.PublishToken(taskID, []byte(tok)); err != nil {
|
||
t.Fatalf("publish token: %v", err)
|
||
}
|
||
}
|
||
if err := dp.CompleteStream(taskID); err != nil {
|
||
t.Fatalf("complete stream: %v", err)
|
||
}
|
||
|
||
select {
|
||
case <-done:
|
||
joined := ""
|
||
for _, s := range got {
|
||
joined += s
|
||
}
|
||
if joined != "Hello Agent!" {
|
||
t.Fatalf("token stream = %q, want %q", joined, "Hello Agent!")
|
||
}
|
||
t.Logf("✓ Token 流闭环:Dispatcher 回流 %d 个 token → Gateway 拼回 %q", len(got), joined)
|
||
case <-time.After(5 * time.Second):
|
||
t.Fatal("timeout: 未收到流结束信号")
|
||
}
|
||
}
|
||
|
||
// TestConcurrentConsume 验证并发消费:一个长时间阻塞的任务(模拟 HITL 待审)
|
||
// 不应阻塞后续任务——旧的串行消费下本测试会超时失败。
|
||
func TestConcurrentConsume(t *testing.T) {
|
||
url := startEmbeddedNATS(t)
|
||
|
||
gw, err := bus.Connect(url)
|
||
if err != nil {
|
||
t.Fatalf("gateway connect: %v", err)
|
||
}
|
||
defer gw.Close()
|
||
|
||
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
|
||
defer cancel()
|
||
if err := gw.EnsureTaskStream(ctx); 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()
|
||
|
||
blockA := make(chan struct{})
|
||
doneB := make(chan string, 1)
|
||
drain, err := dp.ConsumeTasks(ctx, func(_ context.Context, task *contract.Task) error {
|
||
switch task.ID {
|
||
case "A":
|
||
<-blockA // 模拟审批长阻塞
|
||
case "B":
|
||
doneB <- task.ID
|
||
}
|
||
return nil
|
||
})
|
||
if err != nil {
|
||
t.Fatalf("consume: %v", err)
|
||
}
|
||
defer drain(context.Background())
|
||
defer close(blockA)
|
||
|
||
// 先发 A 并等它被消费、卡在 handler 里;再发 B。
|
||
if _, err := gw.PublishTask(ctx, &contract.Task{ID: "A", Graph: json.RawMessage(`{}`)}); err != nil {
|
||
t.Fatalf("publish A: %v", err)
|
||
}
|
||
time.Sleep(400 * time.Millisecond)
|
||
if _, err := gw.PublishTask(ctx, &contract.Task{ID: "B", Graph: json.RawMessage(`{}`)}); err != nil {
|
||
t.Fatalf("publish B: %v", err)
|
||
}
|
||
|
||
select {
|
||
case id := <-doneB:
|
||
if id != "B" {
|
||
t.Fatalf("unexpected task done: %q", id)
|
||
}
|
||
t.Log("✓ 并发消费生效:A 仍阻塞时 B 已完成")
|
||
case <-time.After(4 * time.Second):
|
||
t.Fatal("B 被阻塞的 A 卡住——并发消费未生效(仍是串行)")
|
||
}
|
||
}
|
||
|
||
// TestConcurrentToolServe 验证单实例工具服务并发:一个慢工具调用不阻塞另一个快调用。
|
||
// 旧的串行 ServeTool 下本测试会超时失败。
|
||
func TestConcurrentToolServe(t *testing.T) {
|
||
url := startEmbeddedNATS(t)
|
||
|
||
srv, err := bus.Connect(url)
|
||
if err != nil {
|
||
t.Fatalf("server connect: %v", err)
|
||
}
|
||
defer srv.Close()
|
||
|
||
blockSlow := make(chan struct{})
|
||
drain, err := srv.ServeTool(contract.SubjectToolsGoAll, contract.QueueToolsGo,
|
||
func(_ context.Context, call *contract.ToolCall) *contract.ToolResult {
|
||
if call.Tool == "slow" {
|
||
<-blockSlow // 模拟慢工具(如 RAG 嵌入)
|
||
}
|
||
return &contract.ToolResult{OK: true, Content: call.Tool}
|
||
})
|
||
if err != nil {
|
||
t.Fatalf("serve tool: %v", err)
|
||
}
|
||
defer drain(context.Background())
|
||
defer close(blockSlow)
|
||
|
||
cli, err := bus.Connect(url)
|
||
if err != nil {
|
||
t.Fatalf("client connect: %v", err)
|
||
}
|
||
defer cli.Close()
|
||
|
||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||
defer cancel()
|
||
|
||
// 先发 slow(会阻塞在 handler 里),再发 fast——fast 应在 slow 仍阻塞时返回。
|
||
go func() {
|
||
_, _ = cli.CallTool(ctx, contract.ToolSubjectGo("slow"), &contract.ToolCall{Tool: "slow"})
|
||
}()
|
||
time.Sleep(300 * time.Millisecond)
|
||
|
||
fctx, fcancel := context.WithTimeout(ctx, 3*time.Second)
|
||
defer fcancel()
|
||
res, err := cli.CallTool(fctx, contract.ToolSubjectGo("fast"), &contract.ToolCall{Tool: "fast"})
|
||
if err != nil {
|
||
t.Fatalf("fast 工具被慢工具阻塞——单实例并发未生效: %v", err)
|
||
}
|
||
if res.Content != "fast" {
|
||
t.Fatalf("unexpected result: %q", res.Content)
|
||
}
|
||
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)
|
||
}
|