248 lines
8.0 KiB
Go
248 lines
8.0 KiB
Go
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)
|
||
}
|
||
}
|