package bus_test import ( "context" "encoding/json" "sync/atomic" "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) } // TestGatewayQueueDedup 验证多网关副本下事件不重复处理:两个网关实例都订阅评测/用量/状态, // 经队列组应「每条只被一个副本处理」(旧的广播订阅会被两副本各处理一遍 → 重复落库/重复计费)。 func TestGatewayQueueDedup(t *testing.T) { url := startEmbeddedNATS(t) pub, err := bus.Connect(url) if err != nil { t.Fatalf("pub connect: %v", err) } defer pub.Close() gwA, err := bus.Connect(url) if err != nil { t.Fatalf("gwA connect: %v", err) } defer gwA.Close() gwB, err := bus.Connect(url) if err != nil { t.Fatalf("gwB connect: %v", err) } defer gwB.Close() var total int64 count := func(_ *contract.EvalEvent) { atomic.AddInt64(&total, 1) } if _, err := gwA.SubscribeEval(count); err != nil { t.Fatalf("gwA sub: %v", err) } if _, err := gwB.SubscribeEval(count); err != nil { t.Fatalf("gwB sub: %v", err) } time.Sleep(100 * time.Millisecond) // 等订阅就绪 const n = 50 for i := 0; i < n; i++ { if err := pub.PublishEval(&contract.EvalEvent{TaskID: "t", Overall: 1}); err != nil { t.Fatalf("publish: %v", err) } } time.Sleep(400 * time.Millisecond) // 等投递 got := atomic.LoadInt64(&total) if got != n { t.Fatalf("两副本共处理 %d 条,期望 %d(广播会得 %d=重复处理)", got, n, 2*n) } t.Logf("✓ 队列组去重:%d 条事件被两网关副本合计处理 %d 次(无重复)", n, got) }