server: Go 服务器(REST+WSS 网关+JWT+Argon2+配对+素材直链+任务状态机);修复并发下线 send-on-closed-channel、下发查询 SQL 优先级、配对码原子占用、上传体积上限、JWT 默认密钥告警

This commit is contained in:
Qiufeng
2026-08-20 20:38:21 +08:00
parent bf25d8beac
commit cc8845dcec
44 changed files with 4995 additions and 0 deletions
+247
View File
@@ -0,0 +1,247 @@
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)
}
}