server: Go 服务器(REST+WSS 网关+JWT+Argon2+配对+素材直链+任务状态机);修复并发下线 send-on-closed-channel、下发查询 SQL 优先级、配对码原子占用、上传体积上限、JWT 默认密钥告警

This commit is contained in:
Qiufeng
2026-08-20 20:38:21 +08:00
parent bf25d8beac
commit cc8845dcec
44 changed files with 4995 additions and 0 deletions
+250
View File
@@ -0,0 +1,250 @@
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)
}
+45
View File
@@ -0,0 +1,45 @@
package ws
import (
"log"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"everypublish/server/internal/auth"
"everypublish/server/internal/config"
)
// HandleBrowser 浏览器实时通道:?token=JWT 鉴权,只读订阅工作空间事件
func HandleBrowser(cfg *config.Config, hub *Hub) gin.HandlerFunc {
return func(c *gin.Context) {
claims, err := auth.ParseAccess(cfg.JWTSecret, c.Query("token"))
if err != nil {
c.JSON(401, gin.H{"code": 1002, "message": "令牌无效"})
return
}
conn, err := websocket.Accept(c.Writer, c.Request, &websocket.AcceptOptions{InsecureSkipVerify: true})
if err != nil {
return
}
sess := &Conn{
ID: randHex(16),
Kind: "browser",
WSID: claims.WorkspaceID,
ws: conn,
send: make(chan []byte, 64),
}
hub.RegisterBrowser(sess)
go sess.writeLoop()
log.Printf("[ws] browser session %s online (workspace %d)", sess.ID, claims.WorkspaceID)
defer func() {
hub.UnregisterBrowser(sess)
log.Printf("[ws] browser session %s offline", sess.ID)
}()
for {
if _, _, err = conn.Read(c.Request.Context()); err != nil {
return
}
}
}
}
+161
View File
@@ -0,0 +1,161 @@
package ws
import (
"encoding/json"
"log"
"strconv"
"time"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
"everypublish/server/internal/task"
"everypublish/shared/proto"
)
// Dispatcher 任务下发器(MVP:进程内轮询 + 审批/重试即时触发;接口预留二期换 asynq)
type Dispatcher struct {
DB *gorm.DB
Hub *Hub
Cfg *config.Config
Interval time.Duration
}
// NewDispatcher 构造下发器
func NewDispatcher(db *gorm.DB, hub *Hub, cfg *config.Config) *Dispatcher {
return &Dispatcher{DB: db, Hub: hub, Cfg: cfg, Interval: time.Second}
}
// Start 后台轮询:把到期且排队的任务下发给在线设备
func (d *Dispatcher) Start(stop <-chan struct{}) {
if d.Interval <= 0 {
d.Interval = time.Second
}
ticker := time.NewTicker(d.Interval)
defer ticker.Stop()
for {
select {
case <-stop:
return
case <-ticker.C:
d.poll()
}
}
}
func (d *Dispatcher) poll() {
var rows []models.Task
d.DB.Where("status = ?", task.Queued).
Where("(schedule_at is null or schedule_at <= ?)", time.Now()).
Order("priority desc, id asc").Limit(20).Find(&rows)
for i := range rows {
if d.TryDispatch(rows[i].ID) {
log.Printf("[dispatch] task %d pushed", rows[i].ID)
}
}
}
// TryDispatch 立即下发单个任务(审批/重试后调用);无在线设备返回 false
func (d *Dispatcher) TryDispatch(taskID uint64) bool {
var row models.Task
if err := d.DB.First(&row, taskID).Error; err != nil || row.Status != task.Queued {
return false
}
if row.ScheduleAt != nil && row.ScheduleAt.After(time.Now()) {
return false
}
devID := d.pickDevice(&row)
if devID == 0 {
return false
}
push := d.buildPush(&row, devID)
if push == nil {
return false
}
// 先落库 dispatched,再推送,避免 ack 早于状态落库的竞态
if next, ok := task.Next(row.Status, task.ActionDispatch); ok {
_ = d.DB.Model(&row).Update("status", next).Error
row.Status = next
}
env := proto.NewEnvelope(uuid.NewString(), proto.TypeTaskPush, push)
if err := d.Hub.SendToAgent(devID, env); err != nil {
// 下发失败(设备离线/队列满)回退 queued,避免卡在 dispatched
_ = d.DB.Model(&row).Update("status", task.Queued).Error
return false
}
d.Hub.BroadcastWS(row.WorkspaceID, "task.status", taskEvent(row))
return true
}
// pickDevice 选设备:账号已绑定设备则必须在线上才下发(保持一账号一IP),
// 未绑定设备的账号才兜底任意在线设备。
func (d *Dispatcher) pickDevice(row *models.Task) uint64 {
var accountIDs []uint64
_ = json.Unmarshal([]byte(row.AccountIDs), &accountIDs)
if len(accountIDs) > 0 {
var accounts []models.Account
d.DB.Where("id in ? and workspace_id = ?", accountIDs, row.WorkspaceID).Find(&accounts)
hasAssigned := false
for _, a := range accounts {
if a.AgentDeviceID == 0 {
continue
}
hasAssigned = true
if d.Hub.AgentOnline(a.AgentDeviceID) {
return a.AgentDeviceID
}
}
if hasAssigned {
return 0 // 已绑定设备但离线:不下发,保持 IP 隔离
}
}
online := d.Hub.OnlineDevices(row.WorkspaceID)
if len(online) > 0 {
return online[0]
}
return 0
}
// buildPush 组装 TaskPush(素材生成一次性签名直链)
func (d *Dispatcher) buildPush(row *models.Task, devID uint64) *proto.TaskPush {
var materialIDs []uint64
_ = json.Unmarshal([]byte(row.MaterialIDs), &materialIDs)
urls := make([]string, 0, len(materialIDs))
for _, mid := range materialIDs {
var mat models.Material
if err := d.DB.First(&mat, mid).Error; err != nil {
continue
}
token := models.TransferToken{MaterialID: mat.ID, Token: uuid.NewString(), ExpiresAt: time.Now().Add(10 * time.Minute)}
if err := d.DB.Create(&token).Error; err != nil {
continue
}
urls = append(urls, d.Cfg.BaseURL+"/api/v1/files/"+token.Token)
}
var accountIDs []uint64
_ = json.Unmarshal([]byte(row.AccountIDs), &accountIDs)
var tags []string
_ = json.Unmarshal([]byte(row.Tags), &tags)
push := &proto.TaskPush{
TaskID: strconv.FormatUint(row.ID, 10),
Title: row.Title,
Content: row.Content,
Tags: tags,
MaterialURLs: urls,
Priority: row.Priority,
}
if len(accountIDs) > 0 {
var acc models.Account
if err := d.DB.First(&acc, accountIDs[0]).Error; err == nil {
push.Platform = acc.Platform
push.AccountID = strconv.FormatUint(acc.ID, 10)
push.AccountName = acc.AccountName
}
}
if row.ScheduleAt != nil {
push.ScheduleAt = row.ScheduleAt.UnixMilli()
}
return push
}
+229
View File
@@ -0,0 +1,229 @@
package ws
import (
"encoding/json"
"errors"
"strconv"
"sync"
"time"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"everypublish/shared/proto"
)
// parseID 解析设备ID字符串(失败返回 0)
func parseID(s string) uint64 {
id, _ := strconv.ParseUint(s, 10, 64)
return id
}
// Conn 一条已认证连接(Agent 或浏览器)
type Conn struct {
ID string
Kind string
WSID uint64
ws *websocket.Conn
send chan []byte
closeOnce sync.Once
mu sync.Mutex
closed bool
}
// close 安全关闭发送通道(仅一次)。
// 关闭在 mu 保护下进行;所有发送侧经 sendRaw 先检查 closed,
// 从根本上避免「send on closed channel」panic(关闭与发送的竞态)。
func (c *Conn) close() {
c.closeOnce.Do(func() {
c.mu.Lock()
c.closed = true
close(c.send)
c.mu.Unlock()
})
}
// sendRaw 向发送通道投递:closed 直接返回 false 不恐慌;
// block=true 时队列满将阻塞最多 1s(可靠消息用),否则满即丢弃(广播用)。
func (c *Conn) sendRaw(raw []byte, block bool) bool {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return false
}
if block {
select {
case c.send <- raw:
return true
case <-time.After(time.Second):
return false
}
}
select {
case c.send <- raw:
return true
default:
return false
}
}
// sendEnvelope 统一写出口:所有消息经 send 通道由 writeLoop 单协程写入(避免并发写)
func (c *Conn) sendEnvelope(msg *proto.Envelope) {
raw, err := json.Marshal(msg)
if err != nil {
return
}
c.sendRaw(raw, true)
}
// Hub 连接中枢(goroutine-safe)
type Hub struct {
mu sync.RWMutex
agents map[uint64]*Conn
browser map[string]*Conn
}
// NewHub 创建中枢
func NewHub() *Hub {
return &Hub{
agents: make(map[uint64]*Conn),
browser: make(map[string]*Conn),
}
}
// RegisterAgent 上线登记(同设备重连替换旧连接)
// 网络关闭放到锁外执行,避免阻塞中枢其它操作。
func (h *Hub) RegisterAgent(c *Conn) {
var old *Conn
h.mu.Lock()
if o, ok := h.agents[parseID(c.ID)]; ok {
old = o
}
h.agents[parseID(c.ID)] = c
h.mu.Unlock()
if old != nil {
old.close()
if old.ws != nil {
_ = old.ws.Close(websocket.StatusGoingAway, "replaced by new connection")
}
}
}
// UnregisterAgent 下线(仅当仍是当前连接时才移除,避免旧连接误删新连接)
func (h *Hub) UnregisterAgent(c *Conn) {
var toClose *Conn
h.mu.Lock()
if cur, ok := h.agents[parseID(c.ID)]; ok && cur == c {
delete(h.agents, parseID(c.ID))
toClose = cur
}
h.mu.Unlock()
if toClose != nil {
toClose.close()
}
}
// RegisterBrowser 浏览器连接登记
func (h *Hub) RegisterBrowser(c *Conn) {
h.mu.Lock()
h.browser[c.ID] = c
h.mu.Unlock()
}
// UnregisterBrowser 浏览器断开(指针比较,防串号)
func (h *Hub) UnregisterBrowser(c *Conn) {
var toClose *Conn
h.mu.Lock()
if cur, ok := h.browser[c.ID]; ok && cur == c {
delete(h.browser, c.ID)
toClose = cur
}
h.mu.Unlock()
if toClose != nil {
toClose.close()
}
}
// Clear 清空所有连接(测试/停机用)
func (h *Hub) Clear() {
h.mu.Lock()
for id, c := range h.agents {
c.close()
_ = c.ws.Close(websocket.StatusGoingAway, "hub clear")
delete(h.agents, id)
}
for id, c := range h.browser {
c.close()
_ = c.ws.Close(websocket.StatusGoingAway, "hub clear")
delete(h.browser, id)
}
h.mu.Unlock()
}
// SendToAgent 给设备发消息(1s 超时;队列满丢弃)
func (h *Hub) SendToAgent(deviceID uint64, msg *proto.Envelope) error {
h.mu.RLock()
c, ok := h.agents[deviceID]
h.mu.RUnlock()
if !ok {
return errOffline
}
raw, err := json.Marshal(msg)
if err != nil {
return err
}
if !c.sendRaw(raw, true) {
return errBusy
}
return nil
}
// KickAgent 强制断开设备连接(吊销/远程指令用)
func (h *Hub) KickAgent(deviceID uint64) {
h.mu.RLock()
c, ok := h.agents[deviceID]
h.mu.RUnlock()
if ok {
_ = c.ws.Close(websocket.StatusPolicyViolation, "device revoked")
}
}
// AgentOnline 设备是否在线
func (h *Hub) AgentOnline(deviceID uint64) bool {
h.mu.RLock()
defer h.mu.RUnlock()
_, ok := h.agents[deviceID]
return ok
}
// OnlineDevices 工作空间在线设备列表
func (h *Hub) OnlineDevices(wsid uint64) []uint64 {
h.mu.RLock()
defer h.mu.RUnlock()
ids := make([]uint64, 0)
for id, c := range h.agents {
if c.WSID == wsid {
ids = append(ids, id)
}
}
return ids
}
// BroadcastWS 向工作空间所有浏览器连接广播事件
func (h *Hub) BroadcastWS(wsid uint64, typ string, payload interface{}) {
raw, err := json.Marshal(gin.H{"type": typ, "payload": payload})
if err != nil {
return
}
h.mu.RLock()
defer h.mu.RUnlock()
for _, c := range h.browser {
if c.WSID != wsid {
continue
}
c.sendRaw(raw, false)
}
}
var errOffline = errors.New("device offline")
var errBusy = errors.New("device send queue full")
+44
View File
@@ -0,0 +1,44 @@
package ws
import (
"testing"
)
// TestConnSendAfterClose 验证:连接关闭后 sendRaw/sendEnvelope 不 panic(回归 send-on-closed-channel)。
func TestConnSendAfterClose(t *testing.T) {
c := &Conn{ID: "1", send: make(chan []byte, 1), Kind: "agent"}
c.close()
// 已关闭,发送应返回 false 而非 panic
if c.sendRaw([]byte("x"), true) {
t.Fatal("send to closed conn should return false")
}
if c.sendRaw([]byte("x"), false) {
t.Fatal("send to closed conn (drop) should return false")
}
}
// TestConnDoubleClose 双重 close 不应 panic(closeOnce 保护)
func TestConnDoubleClose(t *testing.T) {
c := &Conn{ID: "1", send: make(chan []byte, 1)}
c.close()
c.close() // 第二次应无效果
}
// TestRegisterReplaceOld 同设备重连替换旧连接:旧连接被 close,映射指向新连接,且不 panic
func TestRegisterReplace(t *testing.T) {
h := NewHub()
old := &Conn{ID: "7", send: make(chan []byte, 64)}
newc := &Conn{ID: "7", send: make(chan []byte, 64)}
h.RegisterAgent(old)
h.RegisterAgent(newc)
if !h.AgentOnline(7) {
t.Fatal("device 7 should be online via new conn")
}
h.UnregisterAgent(old) // 注销旧连接不应影响新连接
if !h.AgentOnline(7) {
t.Fatal("device 7 should still be online after old conn unregister")
}
if old.sendRaw([]byte("x"), false) {
t.Fatal("old conn should be closed")
}
}
+30
View File
@@ -0,0 +1,30 @@
package ws
import (
"crypto/ed25519"
"crypto/rand"
"encoding/hex"
"os"
"path/filepath"
)
// LoadOrCreateServerKey 加载/生成服务器 Ed25519 密钥(hex 文件)
func LoadOrCreateServerKey(path string) (ed25519.PrivateKey, error) {
if raw, err := os.ReadFile(path); err == nil {
seed, err := hex.DecodeString(string(raw))
if err == nil && len(seed) == ed25519.SeedSize {
return ed25519.NewKeyFromSeed(seed), nil
}
}
_, priv, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
return nil, err
}
if err = os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return nil, err
}
if err = os.WriteFile(path, []byte(hex.EncodeToString(priv.Seed())), 0o600); err != nil {
return nil, err
}
return priv, nil
}