208 lines
5.6 KiB
Go
208 lines
5.6 KiB
Go
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")
|
|
}
|
|
}
|