335 lines
9.2 KiB
Go
335 lines
9.2 KiB
Go
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)
|
||
}
|