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) }