251 lines
9.0 KiB
Go
251 lines
9.0 KiB
Go
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)
|
|
}
|