package api_test import ( "context" "crypto/ed25519" "crypto/rand" "encoding/base64" "encoding/json" "fmt" "net/http/httptest" "strconv" "strings" "testing" "time" "github.com/coder/websocket" "github.com/gin-gonic/gin" "everypublish/server/internal/models" "everypublish/shared/proto" ) func readEnv(t *testing.T, conn *websocket.Conn) *proto.Envelope { t.Helper() ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() _, raw, err := conn.Read(ctx) if err != nil { t.Fatalf("ws read: %v", err) } var env proto.Envelope if err = json.Unmarshal(raw, &env); err != nil { t.Fatalf("ws unmarshal: %v", err) } return &env } func writeEnv(t *testing.T, conn *websocket.Conn, env *proto.Envelope) { t.Helper() raw, _ := json.Marshal(env) ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) defer cancel() if err := conn.Write(ctx, websocket.MessageText, raw); err != nil { t.Fatalf("ws write: %v", err) } } func TestWSAgentFullLoop(t *testing.T) { resetTestDB() access, _ := register(t, "agentuser@test.com") // 1. 配对 code, env := doReq(t, "POST", "/api/v1/agent/pair-code", access, nil) if code != 200 || env.Code != 0 { t.Fatalf("pair-code failed: %d %s", code, env.Message) } var pc struct { Code string `json:"code"` } _ = json.Unmarshal(env.Data, &pc) pub, priv, _ := ed25519.GenerateKey(rand.Reader) pubB64 := base64.StdEncoding.EncodeToString(pub) code, env = doReq(t, "POST", "/api/v1/agent/pair", "", gin.H{ "code": pc.Code, "deviceName": "test-device", "os": "test", "version": "0.0.1", "publicKey": pubB64, }) if code != 200 || env.Code != 0 { t.Fatalf("pair failed: %d %s", code, env.Message) } var pd struct { DeviceID uint64 `json:"deviceId"` } _ = json.Unmarshal(env.Data, &pd) if pd.DeviceID == 0 { t.Fatal("deviceId missing") } // 2. 坏签名 hello 应被拒 srv := httptest.NewServer(testRouter) defer srv.Close() wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws/agent" badConn, _, err := websocket.Dial(context.Background(), wsURL, nil) if err != nil { t.Fatalf("dial: %v", err) } badHello := proto.NewEnvelope("id1", proto.TypeHello, proto.DeviceHello{ DeviceID: strconv.FormatUint(pd.DeviceID, 10), Nonce: "n1", TS: time.Now().UnixMilli(), Sig: base64.StdEncoding.EncodeToString([]byte("bad-sig")), Version: "0.0.1", }) rawBad, _ := json.Marshal(badHello) _ = badConn.Write(context.Background(), websocket.MessageText, rawBad) ctx2, cancel2 := context.WithTimeout(context.Background(), 3*time.Second) _, _, err = badConn.Read(ctx2) cancel2() if err == nil { t.Fatal("bad signature should be rejected") } // 3. 正常握手 conn, _, err := websocket.Dial(context.Background(), wsURL, nil) if err != nil { t.Fatalf("dial: %v", err) } nonce := "agent-test-nonce" ts := time.Now().UnixMilli() sig := ed25519.Sign(priv, []byte(nonce+"|"+strconv.FormatInt(ts, 10))) writeEnv(t, conn, proto.NewEnvelope("id2", proto.TypeHello, proto.DeviceHello{ DeviceID: strconv.FormatUint(pd.DeviceID, 10), Nonce: nonce, TS: ts, Sig: base64.StdEncoding.EncodeToString(sig), Version: "0.0.1", })) ackEnv := readEnv(t, conn) if ackEnv.Type != proto.TypeHelloAck { t.Fatalf("expect hello.ack, got %s", ackEnv.Type) } var ack proto.HelloAck _ = json.Unmarshal(ackEnv.Payload, &ack) if ack.SessionID == "" || ack.ServerNonce == "" { t.Fatal("hello.ack missing fields") } // 4. 心跳 → heartbeat.ack writeEnv(t, conn, proto.NewEnvelope("id3", proto.TypeHeartbeat, proto.Heartbeat{DeviceID: "1", TS: time.Now().UnixMilli()})) hbAck := readEnv(t, conn) if hbAck.Type != proto.TypeHeartbeatAck { t.Fatalf("expect heartbeat.ack, got %s", hbAck.Type) } // 5. 账号绑定 → challenge.new → solve → active code, env = doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "douyin", "accountName": "抖音号01"}) if code != 200 { t.Fatalf("create account failed: %d", code) } var acc models.Account _ = json.Unmarshal(env.Data, &acc) code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/accounts/%d/bind", acc.ID), access, nil) if code != 200 { t.Fatalf("bind failed: %d", code) } chEnv := readEnv(t, conn) if chEnv.Type != proto.TypeChallenge { t.Fatalf("expect challenge.new, got %s", chEnv.Type) } var ch proto.Challenge _ = json.Unmarshal(chEnv.Payload, &ch) if ch.AccountID == "" { t.Fatal("challenge missing accountId") } writeEnv(t, conn, proto.NewEnvelope("id5", proto.TypeChallengeSolve, proto.ChallengeSolve{ChallengeID: ch.ChallengeID, Value: "fake-token"})) time.Sleep(200 * time.Millisecond) code, env = doReq(t, "GET", "/api/v1/accounts?status=active", access, nil) if code != 200 { t.Fatalf("accounts list failed: %d", code) } var al struct { List []models.Account `json:"list"` } _ = json.Unmarshal(env.Data, &al) if len(al.List) != 1 || al.List[0].AgentDeviceID != pd.DeviceID { t.Fatalf("account should be active bound to device, got %+v", al.List) } // 6. 任务闭环(含下发延迟测量)+ 浏览器实时事件 browserConn, _, err := websocket.Dial(context.Background(), "ws"+strings.TrimPrefix(srv.URL, "http")+"/ws/browser?token="+access, nil) if err != nil { t.Fatalf("browser dial: %v", err) } matID, _ := uploadMaterial(t, access, "loop.mp4", []byte("loop-content")) code, env = doReq(t, "POST", "/api/v1/tasks", access, gin.H{ "title": "实时闭环任务", "content": "正文", "accountIds": []uint64{acc.ID}, "materialIds": []uint64{matID}, }) if code != 200 { t.Fatalf("create task failed: %d", code) } var taskRow models.Task _ = json.Unmarshal(env.Data, &taskRow) code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/submit", taskRow.ID), access, nil) if code != 200 { t.Fatalf("submit failed: %d", code) } t0 := time.Now() code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/approve", taskRow.ID), access, nil) if code != 200 { t.Fatalf("approve failed: %d", code) } pushEnv := readEnv(t, conn) latency := time.Since(t0) t.Logf("下发延迟 approve→task.push = %s", latency) if latency > 2*time.Second { t.Fatalf("dispatch latency too high: %s", latency) } if pushEnv.Type != proto.TypeTaskPush { t.Fatalf("expect task.push, got %s", pushEnv.Type) } var push proto.TaskPush _ = json.Unmarshal(pushEnv.Payload, &push) if push.TaskID == "" || len(push.MaterialURLs) == 0 { t.Fatalf("task.push missing fields: %+v", push) } writeEnv(t, conn, proto.NewEnvelope(pushEnv.ID, proto.TypeTaskAck, proto.TaskAck{TaskID: push.TaskID, Accept: true})) writeEnv(t, conn, proto.NewEnvelope("id6", proto.TypeTaskResult, proto.TaskResult{ TaskID: push.TaskID, Status: "success", PublishedURL: "https://example.com/v/1", FinishedAt: time.Now().UnixMilli(), })) // 浏览器侧应收到三条 task.status(dispatched/running/success),第三条为 success bEvents := []string{} var lastPayload struct { Status string `json:"status"` } for i := 0; i < 3; i++ { e := readEnv(t, browserConn) bEvents = append(bEvents, e.Type) _ = json.Unmarshal(e.Payload, &lastPayload) } t.Logf("browser events: %v (last=%s)", bEvents, lastPayload.Status) if bEvents[0] != "task.status" || bEvents[1] != "task.status" || bEvents[2] != "task.status" { t.Fatalf("browser should receive 3 task.status events, got %v", bEvents) } if lastPayload.Status != "success" { t.Fatalf("last event should be success, got %s", lastPayload.Status) } // 落库校验(轮询至 success) deadline := time.Now().Add(3 * time.Second) for { code, env = doReq(t, "GET", fmt.Sprintf("/api/v1/tasks/%d", taskRow.ID), access, nil) if code != 200 { t.Fatalf("task get failed: %d", code) } _ = json.Unmarshal(env.Data, &taskRow) if taskRow.Status == "success" || time.Now().After(deadline) { break } time.Sleep(20 * time.Millisecond) } if taskRow.Status != "success" { t.Fatalf("expect success, got %s", taskRow.Status) } if !strings.Contains(taskRow.PublishedURLs, "https://example.com/v/1") { t.Fatalf("publishedUrls not recorded: %s", taskRow.PublishedURLs) } }