Files
EveryPublish/client/core/internal/wsclient/client.go
T

335 lines
9.2 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)
}