client: Windows 客户端(WinUI3 壳 + agent-core;调试用 Electron 壳存档);localserver token 改为常数时间比较
This commit is contained in:
@@ -0,0 +1,334 @@
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/coder/websocket"
|
||||
|
||||
"everypublish/agentcore/internal/device"
|
||||
"everypublish/agentcore/internal/exec"
|
||||
"everypublish/shared/proto"
|
||||
)
|
||||
|
||||
// Client WSS 客户端:指数退避重连 + 任务幂等重放 + 挑战自动处理
|
||||
type Client struct {
|
||||
ServerURL string
|
||||
Identity *device.Identity
|
||||
Registry *exec.Registry
|
||||
ExecTimeout time.Duration
|
||||
|
||||
mu sync.Mutex
|
||||
writeMu sync.Mutex
|
||||
conn *websocket.Conn
|
||||
backoff time.Duration
|
||||
executed map[string]exec.Result
|
||||
|
||||
chMu sync.Mutex
|
||||
challenges map[string]proto.Challenge // 活跃挑战(WinUI 壳展示)
|
||||
online bool
|
||||
sessionID string
|
||||
}
|
||||
|
||||
// New 构造客户端
|
||||
func New(serverURL string, id *device.Identity, registry *exec.Registry) *Client {
|
||||
return &Client{
|
||||
ServerURL: serverURL,
|
||||
Identity: id,
|
||||
Registry: registry,
|
||||
ExecTimeout: 5 * time.Minute,
|
||||
executed: make(map[string]exec.Result),
|
||||
challenges: make(map[string]proto.Challenge),
|
||||
}
|
||||
}
|
||||
|
||||
// Run 主循环:断线按 1s/5s/15s/60s 指数退避重连,握手成功重置
|
||||
func (c *Client) Run(ctx context.Context) error {
|
||||
for {
|
||||
err := c.session(ctx)
|
||||
if ctx.Err() != nil {
|
||||
return nil
|
||||
}
|
||||
c.stepBackoff()
|
||||
log.Printf("[ws] connection lost (%v), reconnect in %s", err, c.backoff)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil
|
||||
case <-time.After(c.backoff):
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) stepBackoff() {
|
||||
switch c.backoff {
|
||||
case 0:
|
||||
c.backoff = time.Second
|
||||
case time.Second:
|
||||
c.backoff = 5 * time.Second
|
||||
case 5 * time.Second:
|
||||
c.backoff = 15 * time.Second
|
||||
default:
|
||||
c.backoff = 60 * time.Second
|
||||
}
|
||||
}
|
||||
|
||||
// session 单次连接会话:握手 → 心跳 → 消息循环
|
||||
func (c *Client) session(ctx context.Context) error {
|
||||
wsURL := "ws" + strings.TrimPrefix(c.ServerURL, "http") + "/ws/agent"
|
||||
conn, _, err := websocket.Dial(ctx, wsURL, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer conn.Close(websocket.StatusNormalClosure, "bye")
|
||||
defer func() {
|
||||
c.mu.Lock()
|
||||
c.online = false
|
||||
c.mu.Unlock()
|
||||
}()
|
||||
c.mu.Lock()
|
||||
c.conn = conn
|
||||
c.mu.Unlock()
|
||||
|
||||
// hello 握手(Ed25519 签名)
|
||||
nonce := randHex(16)
|
||||
ts := time.Now().UnixMilli()
|
||||
sig := c.Identity.Sign([]byte(nonce + "|" + strconv.FormatInt(ts, 10)))
|
||||
hello := proto.NewEnvelope(randHex(8), proto.TypeHello, proto.DeviceHello{
|
||||
DeviceID: strconv.FormatUint(c.Identity.DeviceID, 10),
|
||||
Nonce: nonce,
|
||||
TS: ts,
|
||||
Sig: sig,
|
||||
Version: "0.2.0",
|
||||
})
|
||||
if err = c.send(conn, hello); err != nil {
|
||||
return err
|
||||
}
|
||||
rctx, cancel := context.WithTimeout(ctx, 10*time.Second)
|
||||
_, raw, err := conn.Read(rctx)
|
||||
cancel()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var ackEnv proto.Envelope
|
||||
_ = json.Unmarshal(raw, &ackEnv)
|
||||
if ackEnv.Type != proto.TypeHelloAck {
|
||||
return fmt.Errorf("expect hello.ack, got %s", ackEnv.Type)
|
||||
}
|
||||
var ack proto.HelloAck
|
||||
_ = json.Unmarshal(ackEnv.Payload, &ack)
|
||||
c.mu.Lock()
|
||||
c.backoff = 0 // 握手成功重置退避
|
||||
c.online = true
|
||||
c.sessionID = ack.SessionID
|
||||
c.mu.Unlock()
|
||||
log.Printf("[ws] online, session %s", ack.SessionID)
|
||||
|
||||
// 心跳(15s,服务器 90s 超时检测)
|
||||
hbCtx, hbCancel := context.WithCancel(ctx)
|
||||
defer hbCancel()
|
||||
go c.heartbeatLoop(hbCtx)
|
||||
|
||||
// 消息循环
|
||||
for {
|
||||
typ, raw, err := conn.Read(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if typ != websocket.MessageText {
|
||||
continue
|
||||
}
|
||||
var env proto.Envelope
|
||||
if err = json.Unmarshal(raw, &env); err != nil {
|
||||
continue
|
||||
}
|
||||
c.handle(conn, env)
|
||||
}
|
||||
}
|
||||
|
||||
// handle 分发服务端消息
|
||||
func (c *Client) handle(conn *websocket.Conn, env proto.Envelope) {
|
||||
switch env.Type {
|
||||
case proto.TypeTaskPush:
|
||||
var push proto.TaskPush
|
||||
if err := json.Unmarshal(env.Payload, &push); err != nil {
|
||||
return
|
||||
}
|
||||
c.handleTaskPush(conn, env, push)
|
||||
case proto.TypeChallenge:
|
||||
var ch proto.Challenge
|
||||
if err := json.Unmarshal(env.Payload, &ch); err != nil {
|
||||
return
|
||||
}
|
||||
c.handleChallenge(conn, env, ch)
|
||||
case proto.TypeHeartbeatAck:
|
||||
// 忽略
|
||||
default:
|
||||
log.Printf("[ws] ignore msg %s", env.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// handleTaskPush 任务下发:幂等(已执行过直接重放结果)+ 执行 + 回传
|
||||
func (c *Client) handleTaskPush(conn *websocket.Conn, env proto.Envelope, push proto.TaskPush) {
|
||||
c.mu.Lock()
|
||||
if res, ok := c.executed[push.TaskID]; ok {
|
||||
c.mu.Unlock()
|
||||
log.Printf("[ws] task %s already executed, replay result", push.TaskID)
|
||||
_ = c.send(conn, proto.NewEnvelope(env.ID, proto.TypeTaskAck, proto.TaskAck{TaskID: push.TaskID, Accept: true}))
|
||||
_ = c.send(conn, proto.NewEnvelope(randHex(8), proto.TypeTaskResult, proto.TaskResult{
|
||||
TaskID: res.TaskID, Status: res.Status, PublishedURL: res.PublishedURL,
|
||||
Receipts: res.Receipts, Error: res.Error, FinishedAt: time.Now().UnixMilli(),
|
||||
}))
|
||||
return
|
||||
}
|
||||
c.mu.Unlock()
|
||||
|
||||
log.Printf("[ws] TASK push %s platform=%s title=%s", push.TaskID, push.Platform, push.Title)
|
||||
_ = c.send(conn, proto.NewEnvelope(env.ID, proto.TypeTaskAck, proto.TaskAck{TaskID: push.TaskID, Accept: true}))
|
||||
|
||||
task := exec.Task{
|
||||
TaskID: push.TaskID, Platform: push.Platform, AccountID: push.AccountID, AccountName: push.AccountName,
|
||||
Title: push.Title, Content: push.Content, Tags: push.Tags, MaterialURLs: push.MaterialURLs,
|
||||
ScheduleAt: push.ScheduleAt, Priority: push.Priority,
|
||||
}
|
||||
var res exec.Result
|
||||
e := c.Registry.Get(push.Platform)
|
||||
if e == nil {
|
||||
res = exec.Result{TaskID: push.TaskID, Status: "failed", Error: "平台执行器未接入: " + push.Platform}
|
||||
} else {
|
||||
eCtx, cancel := context.WithTimeout(context.Background(), c.ExecTimeout)
|
||||
res = e.Execute(eCtx, task)
|
||||
cancel()
|
||||
}
|
||||
c.mu.Lock()
|
||||
c.executed[push.TaskID] = res
|
||||
if len(c.executed) > 1024 {
|
||||
c.executed = make(map[string]exec.Result)
|
||||
}
|
||||
c.mu.Unlock()
|
||||
_ = c.send(conn, proto.NewEnvelope(randHex(8), proto.TypeTaskResult, proto.TaskResult{
|
||||
TaskID: res.TaskID, Status: res.Status, PublishedURL: res.PublishedURL,
|
||||
Receipts: res.Receipts, Error: res.Error, FinishedAt: time.Now().UnixMilli(),
|
||||
}))
|
||||
log.Printf("[ws] TASK %s done status=%s", push.TaskID, res.Status)
|
||||
}
|
||||
|
||||
// handleChallenge 挑战:ack + 有自动解决器则异步解决(60s 超时),否则留人工
|
||||
func (c *Client) handleChallenge(conn *websocket.Conn, env proto.Envelope, ch proto.Challenge) {
|
||||
log.Printf("[ws] CHALLENGE %s kind=%s prompt=%s", ch.ChallengeID, ch.Kind, ch.Prompt)
|
||||
c.chMu.Lock()
|
||||
c.challenges[ch.ChallengeID] = ch
|
||||
c.chMu.Unlock()
|
||||
_ = c.send(conn, proto.NewEnvelope(env.ID, proto.TypeChallengeAck, proto.ChallengeAck{ChallengeID: ch.ChallengeID, Action: "accept"}))
|
||||
solver := c.Registry.Solve(ch.Platform)
|
||||
if solver == nil {
|
||||
log.Printf("[ws] challenge %s 无自动解决器,等待人工处理", ch.ChallengeID)
|
||||
return
|
||||
}
|
||||
// 平台生成二维码时更新挑战记录(壳经 /challenges 展示)
|
||||
onQR := func(qrToken, qrURL string) {
|
||||
c.chMu.Lock()
|
||||
if cur, ok := c.challenges[ch.ChallengeID]; ok {
|
||||
cur.QRToken = qrToken
|
||||
cur.QRURL = qrURL
|
||||
c.challenges[ch.ChallengeID] = cur
|
||||
}
|
||||
c.chMu.Unlock()
|
||||
}
|
||||
go func() {
|
||||
sCtx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
|
||||
defer cancel()
|
||||
if solved, value := solver.SolveChallenge(sCtx, ch, onQR); solved {
|
||||
_ = c.send(conn, proto.NewEnvelope(randHex(8), proto.TypeChallengeSolve, proto.ChallengeSolve{ChallengeID: ch.ChallengeID, Value: value}))
|
||||
c.chMu.Lock()
|
||||
delete(c.challenges, ch.ChallengeID)
|
||||
c.chMu.Unlock()
|
||||
log.Printf("[ws] CHALLENGE %s solved", ch.ChallengeID)
|
||||
} else {
|
||||
log.Printf("[ws] CHALLENGE %s 未解决,等待人工", ch.ChallengeID)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// heartbeatLoop 心跳(15s)
|
||||
func (c *Client) heartbeatLoop(ctx context.Context) {
|
||||
ticker := time.NewTicker(15 * time.Second)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
c.mu.Lock()
|
||||
conn := c.conn
|
||||
c.mu.Unlock()
|
||||
if conn == nil {
|
||||
continue
|
||||
}
|
||||
_ = c.send(conn, proto.NewEnvelope(randHex(8), proto.TypeHeartbeat, proto.Heartbeat{DeviceID: strconv.FormatUint(c.Identity.DeviceID, 10), TS: time.Now().UnixMilli()}))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// send 统一写出口(写锁串行化,避免并发写)
|
||||
func (c *Client) send(conn *websocket.Conn, msg *proto.Envelope) error {
|
||||
raw, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
wctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
c.writeMu.Lock()
|
||||
defer c.writeMu.Unlock()
|
||||
return conn.Write(wctx, websocket.MessageText, raw)
|
||||
}
|
||||
|
||||
// Challenges 当前活跃挑战快照(WinUI 壳轮询展示)
|
||||
func (c *Client) Challenges() []proto.Challenge {
|
||||
c.chMu.Lock()
|
||||
defer c.chMu.Unlock()
|
||||
list := make([]proto.Challenge, 0, len(c.challenges))
|
||||
for _, ch := range c.challenges {
|
||||
list = append(list, ch)
|
||||
}
|
||||
return list
|
||||
}
|
||||
|
||||
// SolveChallenge 人工完成挑战(验证码/确认),经 WSS 回传服务器
|
||||
func (c *Client) SolveChallenge(challengeID, value string) error {
|
||||
c.mu.Lock()
|
||||
conn := c.conn
|
||||
c.mu.Unlock()
|
||||
if conn == nil {
|
||||
return errors.New("offline")
|
||||
}
|
||||
return c.send(conn, proto.NewEnvelope(randHex(8), proto.TypeChallengeSolve, proto.ChallengeSolve{ChallengeID: challengeID, Value: value}))
|
||||
}
|
||||
|
||||
// Online 当前是否在线
|
||||
func (c *Client) Online() bool {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return c.online
|
||||
}
|
||||
|
||||
// SessionID 当前会话 ID
|
||||
func (c *Client) SessionID() string {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return c.sessionID
|
||||
}
|
||||
|
||||
func randHex(n int) string {
|
||||
b := make([]byte, n)
|
||||
_, _ = rand.Read(b)
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
@@ -0,0 +1,207 @@
|
||||
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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user