client: Windows 客户端(WinUI3 壳 + agent-core;调试用 Electron 壳存档);localserver token 改为常数时间比较

This commit is contained in:
Qiufeng
2026-08-20 20:38:21 +08:00
parent 33b7d57498
commit 9abf7f8213
60 changed files with 10025 additions and 0 deletions
+334
View File
@@ -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")
}
}