package wsclient import ( "context" "crypto/ed25519" "encoding/base64" "encoding/json" "net/http" "net/http/httptest" "strconv" "sync/atomic" "testing" "time" "github.com/coder/websocket" "everypublish/agentcore/internal/device" "everypublish/agentcore/internal/exec" "everypublish/shared/proto" ) // countingExec 计数执行器(验证幂等:只执行一次) type countingExec struct{ n int32 } func (c *countingExec) Platform() string { return "" } func (c *countingExec) Execute(ctx context.Context, task exec.Task) exec.Result { atomic.AddInt32(&c.n, 1) return exec.Result{TaskID: task.TaskID, Status: "success", PublishedURL: "https://ex.com/" + task.TaskID} } func (c *countingExec) SolveChallenge(ctx context.Context, ch proto.Challenge) (bool, string) { return true, "tok" } // verifyHello 校验客户端 hello 签名 func verifyHello(t *testing.T, id *device.Identity, raw []byte) { t.Helper() var env proto.Envelope _ = json.Unmarshal(raw, &env) var hello proto.DeviceHello _ = json.Unmarshal(env.Payload, &hello) pub, _ := base64.StdEncoding.DecodeString(id.PublicKey) sig, _ := base64.StdEncoding.DecodeString(hello.Sig) if !ed25519.Verify(ed25519.PublicKey(pub), []byte(hello.Nonce+"|"+strconv.FormatInt(hello.TS, 10)), sig) { t.Errorf("hello signature invalid") } } // sendAck 回 hello.ack func sendAck(r *http.Request, conn *websocket.Conn) { ack := proto.NewEnvelope("ack1", proto.TypeHelloAck, proto.HelloAck{ServerNonce: "n", SessionID: "s1", ServerTS: 1}) rawAck, _ := json.Marshal(ack) _ = conn.Write(r.Context(), websocket.MessageText, rawAck) } // startStub 启动 stub:验签 → ack → 推任务1 → 收结果 → 再推任务1(重放)→ 收重放结果 func startStub(t *testing.T, id *device.Identity) (*httptest.Server, chan proto.TaskResult, chan struct{}) { results := make(chan proto.TaskResult, 8) connected := make(chan struct{}, 8) srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := websocket.Accept(w, r, nil) if err != nil { return } defer conn.Close(websocket.StatusNormalClosure, "") ctx, cancel := context.WithTimeout(r.Context(), 5*time.Second) _, raw, err := conn.Read(ctx) cancel() if err != nil { return } verifyHello(t, id, raw) sendAck(r, conn) connected <- struct{}{} push := proto.NewEnvelope("p1", proto.TypeTaskPush, proto.TaskPush{TaskID: "1", Platform: "douyin", Title: "t1"}) rawPush, _ := json.Marshal(push) _ = conn.Write(r.Context(), websocket.MessageText, rawPush) gotFirst := false for { _, raw, err = conn.Read(r.Context()) if err != nil { return } var env2 proto.Envelope if json.Unmarshal(raw, &env2) != nil { continue } if env2.Type == proto.TypeTaskResult { var res proto.TaskResult _ = json.Unmarshal(env2.Payload, &res) results <- res if !gotFirst { gotFirst = true push2 := proto.NewEnvelope("p2", proto.TypeTaskPush, proto.TaskPush{TaskID: "1", Platform: "douyin", Title: "t1"}) rawPush2, _ := json.Marshal(push2) _ = conn.Write(r.Context(), websocket.MessageText, rawPush2) } } } })) return srv, results, connected } func TestClientTaskRoundTripAndIdempotent(t *testing.T) { id, _ := device.LoadOrCreate(t.TempDir(), "dev") id.DeviceID = 1 registry := exec.NewRegistry() ce := &countingExec{} registry.SetFallback(ce) srv, results, connected := startStub(t, id) defer srv.Close() c := New(srv.URL, id, registry) ctx, cancel := context.WithCancel(context.Background()) defer cancel() go func() { _ = c.Run(ctx) }() select { case <-connected: case <-time.After(5 * time.Second): t.Fatal("handshake timeout") } var r1 proto.TaskResult select { case r1 = <-results: case <-time.After(5 * time.Second): t.Fatal("first result timeout") } if r1.Status != "success" { t.Fatalf("first result should be success, got %s", r1.Status) } var r2 proto.TaskResult select { case r2 = <-results: case <-time.After(5 * time.Second): t.Fatal("replay result timeout") } if r2.Status != "success" || r2.PublishedURL != r1.PublishedURL { t.Fatalf("replay should return cached result, got %+v", r2) } if n := atomic.LoadInt32(&ce.n); n != 1 { t.Fatalf("executor should run exactly once, got %d", n) } } // TestClientReconnect 服务器断开后客户端应按退避自动重连 func TestClientReconnect(t *testing.T) { id, _ := device.LoadOrCreate(t.TempDir(), "dev") id.DeviceID = 1 registry := exec.NewRegistry() registry.SetFallback(&countingExec{}) var connCount int32 connected := make(chan struct{}, 8) srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { conn, err := websocket.Accept(w, r, nil) if err != nil { return } ctx, cancel := context.WithTimeout(r.Context(), 5*time.Second) _, raw, err := conn.Read(ctx) cancel() if err != nil { return } verifyHello(t, id, raw) sendAck(r, conn) connected <- struct{}{} n := atomic.AddInt32(&connCount, 1) if n == 1 { // 第一次连接:ack 后立刻断开,模拟服务器故障 _ = conn.Close(websocket.StatusGoingAway, "restart") return } // 后续连接:正常读循环 for { if _, _, err = conn.Read(r.Context()); err != nil { return } } })) defer srv.Close() c := New(srv.URL, id, registry) ctx, cancel := context.WithCancel(context.Background()) defer cancel() go func() { _ = c.Run(ctx) }() select { case <-connected: case <-time.After(5 * time.Second): t.Fatal("first connect timeout") } select { case <-connected: case <-time.After(8 * time.Second): t.Fatal("reconnect timeout") } if atomic.LoadInt32(&connCount) < 2 { t.Fatal("expected second connection") } }