package ws import ( "context" "crypto/ed25519" "crypto/rand" "encoding/base64" "encoding/hex" "encoding/json" "log" "strconv" "time" "github.com/coder/websocket" "github.com/gin-gonic/gin" "gorm.io/gorm" "everypublish/server/internal/config" "everypublish/server/internal/models" "everypublish/server/internal/task" "everypublish/shared/proto" ) // signPayload 签名原文:nonce + "|" + ts(两端一致) func signPayload(nonce string, ts int64) []byte { return []byte(nonce + "|" + strconv.FormatInt(ts, 10)) } // HandleAgent 处理 Agent WSS 连接:hello 验签 → 注册 → 消息循环 func HandleAgent(cfg *config.Config, db *gorm.DB, hub *Hub, serverKey ed25519.PrivateKey) gin.HandlerFunc { return func(c *gin.Context) { conn, err := websocket.Accept(c.Writer, c.Request, &websocket.AcceptOptions{InsecureSkipVerify: true}) if err != nil { return } // 首条消息必须是 hello(10s 超时) ctx, cancel := context.WithTimeout(c.Request.Context(), 10*time.Second) defer cancel() typ, raw, err := conn.Read(ctx) if err != nil { _ = conn.Close(websocket.StatusPolicyViolation, "hello timeout") return } if typ != websocket.MessageText { _ = conn.Close(websocket.StatusPolicyViolation, "expect text") return } var helloEnv proto.Envelope if err = json.Unmarshal(raw, &helloEnv); err != nil || helloEnv.Type != proto.TypeHello { _ = conn.Close(websocket.StatusPolicyViolation, "expect hello") return } var hello proto.DeviceHello if err = json.Unmarshal(helloEnv.Payload, &hello); err != nil { _ = conn.Close(websocket.StatusPolicyViolation, "bad hello") return } deviceID, _ := strconv.ParseUint(hello.DeviceID, 10, 64) var device models.AgentDevice if err = db.First(&device, deviceID).Error; err != nil { _ = conn.Close(websocket.StatusPolicyViolation, "unknown device") return } if device.Revoked { _ = conn.Close(websocket.StatusPolicyViolation, "device revoked") return } pubRaw, _ := base64.StdEncoding.DecodeString(device.PublicKey) if len(pubRaw) != ed25519.PublicKeySize { _ = conn.Close(websocket.StatusPolicyViolation, "bad key") return } sig, _ := base64.StdEncoding.DecodeString(hello.Sig) if !ed25519.Verify(ed25519.PublicKey(pubRaw), signPayload(hello.Nonce, hello.TS), sig) { _ = conn.Close(websocket.StatusPolicyViolation, "bad signature") return } // 注册 + 应答 serverNonce := randHex(16) serverTS := time.Now().UnixMilli() ackSig := ed25519.Sign(serverKey, signPayload(serverNonce, serverTS)) sess := &Conn{ ID: hello.DeviceID, Kind: "agent", WSID: device.WorkspaceID, ws: conn, send: make(chan []byte, 64), } hub.RegisterAgent(sess) now := time.Now() _ = db.Model(&device).Updates(map[string]interface{}{"status": "online", "last_seen_at": &now, "ip": c.ClientIP(), "version": hello.Version}).Error hub.BroadcastWS(device.WorkspaceID, "agent.status", gin.H{"deviceId": device.ID, "online": true}) ack := proto.NewEnvelope(randHex(8), proto.TypeHelloAck, proto.HelloAck{ ServerNonce: serverNonce, SessionID: randHex(16), ServerTS: serverTS, Sig: base64.StdEncoding.EncodeToString(ackSig), }) go sess.writeLoop() sess.sendEnvelope(ack) log.Printf("[ws] agent %d online (workspace %d)", device.ID, device.WorkspaceID) // 消息循环 defer func() { hub.UnregisterAgent(sess) _ = db.Model(&device).Updates(map[string]interface{}{"status": "offline", "last_seen_at": time.Now()}).Error hub.BroadcastWS(device.WorkspaceID, "agent.status", gin.H{"deviceId": device.ID, "online": false}) log.Printf("[ws] agent %d offline", device.ID) }() for { // 90s 无消息视为死连接(客户端心跳 10-30s) readCtx, readCancel := context.WithTimeout(c.Request.Context(), 90*time.Second) typ, raw, err = conn.Read(readCtx) readCancel() if err != nil { return } if typ != websocket.MessageText { continue } var env proto.Envelope if err = json.Unmarshal(raw, &env); err != nil { continue } handleAgentMessage(db, hub, sess, device, env) } } } // handleAgentMessage 分发 Agent 消息(所有回写经 sess 统一通道,避免并发写) func handleAgentMessage(db *gorm.DB, hub *Hub, sess *Conn, device models.AgentDevice, env proto.Envelope) { var err error switch env.Type { case proto.TypeHeartbeat: var hb proto.Heartbeat if err = json.Unmarshal(env.Payload, &hb); err == nil { now := time.Now() _ = db.Model(&models.AgentDevice{}).Where("id = ?", device.ID).Update("last_seen_at", &now).Error ack := proto.NewEnvelope(env.ID, proto.TypeHeartbeatAck, proto.HeartbeatAck{ServerTS: now.UnixMilli()}) sess.sendEnvelope(ack) } case proto.TypeTaskAck: var ack proto.TaskAck if err = json.Unmarshal(env.Payload, &ack); err == nil { applyTaskAck(db, hub, device, ack) } case proto.TypeTaskResult: var result proto.TaskResult if err = json.Unmarshal(env.Payload, &result); err == nil { log.Printf("[ws] task.result received: %+v", result) applyTaskResult(db, hub, device, result) } else { log.Printf("[ws] task.result unmarshal error: %v payload=%s", err, string(env.Payload)) } case proto.TypeChallengeAck: var cack proto.ChallengeAck if err = json.Unmarshal(env.Payload, &cack); err == nil { log.Printf("[ws] challenge %s ack: %s", cack.ChallengeID, cack.Action) } case proto.TypeChallengeSolve: var cs proto.ChallengeSolve if err = json.Unmarshal(env.Payload, &cs); err == nil { applyChallengeSolve(db, hub, device, cs) } default: log.Printf("[ws] unknown message type: %s", env.Type) } } // applyTaskAck 任务确认:dispatched → running func applyTaskAck(db *gorm.DB, hub *Hub, device models.AgentDevice, ack proto.TaskAck) { tid, _ := strconv.ParseUint(ack.TaskID, 10, 64) var row models.Task if err := db.Where("id = ? and workspace_id = ?", tid, device.WorkspaceID).First(&row).Error; err != nil { return } if !ack.Accept { _ = db.Model(&row).Updates(map[string]interface{}{"status": task.Failed, "error_message": ack.Reason}).Error hub.BroadcastWS(device.WorkspaceID, "task.status", taskEvent(row)) return } if next, ok := task.Next(row.Status, task.ActionStart); ok { _ = db.Model(&row).Update("status", next).Error row.Status = next hub.BroadcastWS(device.WorkspaceID, "task.status", taskEvent(row)) } } // applyTaskResult 结果回传:running → success/failed func applyTaskResult(db *gorm.DB, hub *Hub, device models.AgentDevice, result proto.TaskResult) { tid, _ := strconv.ParseUint(result.TaskID, 10, 64) var row models.Task if err := db.Where("id = ? and workspace_id = ?", tid, device.WorkspaceID).First(&row).Error; err != nil { log.Printf("[ws] task.result: task %d not found: %v", tid, err) return } log.Printf("[ws] task.result: task %d current status %s, want %s", tid, row.Status, result.Status) if result.Status == "success" { if next, ok := task.Next(row.Status, task.ActionSuccess); ok { urls, _ := json.Marshal([]string{result.PublishedURL}) receipts, _ := json.Marshal(result.Receipts) _ = db.Model(&row).Updates(map[string]interface{}{ "status": next, "published_urls": string(urls), "receipts": string(receipts), "error_message": "", "published_at": time.Now(), }).Error row.Status = next hub.BroadcastWS(device.WorkspaceID, "task.status", taskEvent(row)) log.Printf("[ws] task %d success in workspace %d", tid, device.WorkspaceID) } } else { if next, ok := task.Next(row.Status, task.ActionFail); ok { _ = db.Model(&row).Updates(map[string]interface{}{"status": next, "error_message": result.Error}).Error row.Status = next hub.BroadcastWS(device.WorkspaceID, "task.status", taskEvent(row)) } } } // applyChallengeSolve 挑战完成:账号激活 func applyChallengeSolve(db *gorm.DB, hub *Hub, device models.AgentDevice, cs proto.ChallengeSolve) { cid, _ := strconv.ParseUint(cs.ChallengeID, 10, 64) var ch models.Challenge if err := db.Where("id = ? and workspace_id = ?", cid, device.WorkspaceID).First(&ch).Error; err != nil { return } now := time.Now() _ = db.Model(&ch).Updates(map[string]interface{}{"status": "solved", "solved_at": &now, "qr_token": cs.Value}).Error _ = db.Model(&models.Account{}).Where("id = ?", ch.AccountID).Updates(map[string]interface{}{"status": "active", "agent_device_id": device.ID, "last_active_at": &now}).Error hub.BroadcastWS(device.WorkspaceID, "challenge.status", gin.H{"challengeId": ch.ID, "status": "solved", "accountId": ch.AccountID}) log.Printf("[ws] challenge %d solved, account %d active", ch.ID, ch.AccountID) } // taskEvent 任务状态事件 func taskEvent(row models.Task) gin.H { return gin.H{"taskId": row.ID, "status": row.Status, "errorMessage": row.ErrorMessage, "publishedUrls": row.PublishedURLs} } // writeLoop 写协程(唯一写出口) func (c *Conn) writeLoop() { for raw := range c.send { wctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) err := c.ws.Write(wctx, websocket.MessageText, raw) cancel() if err != nil { return } } } func randHex(n int) string { b := make([]byte, n) _, _ = rand.Read(b) return hex.EncodeToString(b) }