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
+221
View File
@@ -0,0 +1,221 @@
package api_test
import (
"bytes"
"encoding/json"
"fmt"
"mime/multipart"
"net/http/httptest"
"net/textproto"
"strings"
"testing"
"github.com/gin-gonic/gin"
"everypublish/server/internal/models"
)
func multipartBody(t *testing.T, filename, contentType string, content []byte) *bytes.Buffer {
t.Helper()
buf := &bytes.Buffer{}
w := multipart.NewWriter(buf)
h := make(textproto.MIMEHeader)
h.Set("Content-Disposition", fmt.Sprintf(`form-data; name="file"; filename="%s"`, filename))
h.Set("Content-Type", contentType)
part, err := w.CreatePart(h)
if err != nil {
t.Fatalf("create part: %v", err)
}
_, _ = part.Write(content)
_ = w.Close()
return buf
}
func uploadMaterial(t *testing.T, access string, filename string, content []byte) (uint64, string) {
t.Helper()
buf := multipartBody(t, filename, "video/mp4", content)
req := httptest.NewRequest("POST", "/api/v1/materials", buf)
req.Header.Set("Content-Type", "multipart/form-data; boundary="+strings.Split(buf.String(), "\r\n")[0][2:])
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
testRouter.ServeHTTP(w, req)
var env envelope
_ = json.Unmarshal(w.Body.Bytes(), &env)
if w.Code != 200 || env.Code != 0 {
t.Fatalf("upload failed: %d %s", w.Code, env.Message)
}
var data struct {
Material models.Material `json:"material"`
Dedup bool `json:"dedup"`
}
_ = json.Unmarshal(env.Data, &data)
if data.Dedup {
t.Fatal("first upload should not be dedup")
}
return data.Material.ID, data.Material.SHA256
}
func TestAccountChallengeFlow(t *testing.T) {
resetTestDB()
access, _ := register(t, "ops@test.com")
code, env := doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "douyin", "accountName": "抖音测试号"})
if code != 200 || env.Code != 0 {
t.Fatalf("create account failed: %d %s", code, env.Message)
}
var acc models.Account
_ = json.Unmarshal(env.Data, &acc)
// 不支持的平台
code, _ = doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "tiktok", "accountName": "x"})
if code != 400 {
t.Fatalf("unsupported platform should 400, got %d", code)
}
// 绑定 → 挑战
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/accounts/%d/bind", acc.ID), access, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("bind failed: %d %s", code, env.Message)
}
// 绑定中账号不可删
code, _ = doReq(t, "DELETE", fmt.Sprintf("/api/v1/accounts/%d", acc.ID), access, nil)
if code != 409 {
t.Fatalf("delete binding account should 409, got %d", code)
}
// 挑战列表
code, env = doReq(t, "GET", "/api/v1/challenges", access, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("challenge list failed: %d %s", code, env.Message)
}
var cl struct {
List []models.Challenge `json:"list"`
}
_ = json.Unmarshal(env.Data, &cl)
if len(cl.List) != 1 || cl.List[0].Kind != "pending" {
t.Fatalf("expect 1 pending challenge, got %d", len(cl.List))
}
// 挂起 → 重发
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/challenges/%d/suspend", cl.List[0].ID), access, nil)
if code != 200 {
t.Fatalf("suspend failed: %d", code)
}
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/challenges/%d/resend", cl.List[0].ID), access, nil)
if code != 200 {
t.Fatalf("resend failed: %d", code)
}
}
func TestMaterialAndTaskFlow(t *testing.T) {
resetTestDB()
access, _ := register(t, "pub@test.com")
content := []byte("fake-video-bytes-for-sha256-test-0001")
matID, sha := uploadMaterial(t, access, "demo.mp4", content)
if len(sha) != 64 {
t.Fatalf("bad sha256: %s", sha)
}
// 重复上传 → 去重
buf := multipartBody(t, "demo2.mp4", "video/mp4", content)
req := httptest.NewRequest("POST", "/api/v1/materials", buf)
req.Header.Set("Content-Type", "multipart/form-data; boundary="+strings.Split(buf.String(), "\r\n")[0][2:])
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
testRouter.ServeHTTP(w, req)
var env envelope
_ = json.Unmarshal(w.Body.Bytes(), &env)
var data struct {
Material models.Material `json:"material"`
Dedup bool `json:"dedup"`
}
_ = json.Unmarshal(env.Data, &data)
if !data.Dedup || data.Material.ID != matID {
t.Fatalf("dedup expect same id %d, got dedup=%v id=%d", matID, data.Dedup, data.Material.ID)
}
// 签名直链(一次性)
code, env := doReq(t, "GET", fmt.Sprintf("/api/v1/materials/%d/url", matID), access, nil)
if code != 200 {
t.Fatalf("url failed: %d", code)
}
var urld struct {
URL string `json:"url"`
}
_ = json.Unmarshal(env.Data, &urld)
path := strings.TrimPrefix(urld.URL, "http://127.0.0.1:8090")
if path == urld.URL {
t.Fatalf("bad url: %s", urld.URL)
}
code, _ = doReq(t, "GET", path, "", nil)
if code != 200 {
t.Fatalf("file fetch should 200, got %d", code)
}
code, _ = doReq(t, "GET", path, "", nil)
if code != 404 {
t.Fatalf("one-time token should 404 on reuse, got %d", code)
}
// 建任务 → 提交 → 驳回 → 重提 → 通过 → 取消
code, env = doReq(t, "POST", "/api/v1/tasks", access, gin.H{
"title": "新品发布", "content": "正文", "tags": []string{"新品"},
"accountIds": []uint64{1}, "materialIds": []uint64{matID},
})
if code != 200 || env.Code != 0 {
t.Fatalf("create task failed: %d %s", code, env.Message)
}
var taskRow models.Task
_ = json.Unmarshal(env.Data, &taskRow)
if taskRow.Status != "draft" {
t.Fatalf("new task should be draft, got %s", taskRow.Status)
}
// 非法转移:draft 直接 approve 应 409
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/approve", taskRow.ID), access, nil)
if code != 409 {
t.Fatalf("draft approve should 409, got %d", code)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/submit", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("submit failed: %d", code)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/reject", taskRow.ID), access, gin.H{"note": "标题需修改"})
if code != 200 {
t.Fatalf("reject failed: %d", code)
}
code, env = doReq(t, "GET", fmt.Sprintf("/api/v1/tasks/%d", taskRow.ID), access, nil)
_ = json.Unmarshal(env.Data, &taskRow)
if taskRow.Status != "rejected" {
t.Fatalf("expect rejected, got %s", taskRow.Status)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/resubmit", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("resubmit failed: %d", code)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/approve", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("approve failed: %d", code)
}
code, env = doReq(t, "GET", fmt.Sprintf("/api/v1/tasks/%d", taskRow.ID), access, nil)
_ = json.Unmarshal(env.Data, &taskRow)
if taskRow.Status != "queued" {
t.Fatalf("expect queued, got %s", taskRow.Status)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/cancel", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("cancel failed: %d", code)
}
// 通知:驳回+通过 2 条
code, env = doReq(t, "GET", "/api/v1/notifications", access, nil)
if code != 200 {
t.Fatalf("notifications failed: %d", code)
}
var nl struct {
List []models.Notification `json:"list"`
Unread int64 `json:"unread"`
}
_ = json.Unmarshal(env.Data, &nl)
if nl.Unread < 2 {
t.Fatalf("expect >=2 unread, got %d", nl.Unread)
}
code, _ = doReq(t, "POST", "/api/v1/notifications/read-all", access, nil)
if code != 200 {
t.Fatalf("read-all failed: %d", code)
}
code, env = doReq(t, "GET", "/api/v1/notifications", access, nil)
_ = json.Unmarshal(env.Data, &nl)
if nl.Unread != 0 {
t.Fatalf("expect 0 unread after read-all, got %d", nl.Unread)
}
}
+195
View File
@@ -0,0 +1,195 @@
package handlers
import (
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
"everypublish/server/internal/ws"
"everypublish/shared/proto"
)
// AccountHandler 账号台账处理器(凭据仅在 Agent 本机保险库,服务器零凭据)
type AccountHandler struct {
DB *gorm.DB
Hub *ws.Hub
}
var platforms = map[string]bool{
"douyin": true, "kuaishou": true, "xiaohongshu": true, "bilibili": true,
"shipinhao": true, "x": true, "instagram": true, "whatsapp": true, "youtube": true,
}
type accountReq struct {
Platform string `json:"platform" binding:"required"`
Remark string `json:"remark"`
AgentDeviceID uint64 `json:"agentDeviceId"`
AccountName string `json:"accountName"`
}
// List 账号台账(筛选 platform/status)
func (h *AccountHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.Account{}).Where("workspace_id = ?", c.GetUint64("wsid"))
if p := c.Query("platform"); p != "" {
q = q.Where("platform = ?", p)
}
if s := c.Query("status"); s != "" {
q = q.Where("status = ?", s)
}
var total int64
q.Count(&total)
var list []models.Account
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&list)
response.OK(c, gin.H{"list": list, "total": total, "page": page, "size": size})
}
// Create 新增账号(初始 unbound;凭据只登记元数据)
func (h *AccountHandler) Create(c *gin.Context) {
var req accountReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
if !platforms[req.Platform] {
response.Fail(c, http.StatusBadRequest, 1001, "暂不支持该平台")
return
}
account := models.Account{
WorkspaceID: c.GetUint64("wsid"),
Platform: req.Platform,
Remark: req.Remark,
AgentDeviceID: req.AgentDeviceID,
Status: "unbound",
Health: 100,
}
if err := h.DB.Create(&account).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "创建失败")
return
}
response.Audit(c, "account.create", "account:"+itoa(account.ID), gin.H{"platform": req.Platform, "remark": req.Remark})
response.OK(c, account)
}
// Update 修改账号资料(名称/头像/IP画像)
func (h *AccountHandler) Update(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req accountReq
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
updates := map[string]interface{}{"remark": req.Remark}
if req.AccountName != "" {
updates["account_name"] = req.AccountName
}
if req.AgentDeviceID > 0 {
updates["agent_device_id"] = req.AgentDeviceID
}
if err = h.DB.Model(&models.Account{}).Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).Updates(updates).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "account.update", "account:"+itoa(id), gin.H{"remark": req.Remark})
response.OK(c, nil)
}
// Delete 删除账号(仅未绑定)
func (h *AccountHandler) Delete(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var account models.Account
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&account).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "账号不存在")
return
}
if account.Status == "active" || account.Status == "binding" {
response.Fail(c, http.StatusConflict, 3001, "已绑定或绑定中的账号不可删除")
return
}
if err = h.DB.Delete(&account).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "删除失败")
return
}
response.Audit(c, "account.delete", "account:"+itoa(id), nil)
response.OK(c, nil)
}
// Bind 发起绑定:创建挑战记录(pending),等待客户端扫码(D4 WSS 联动)
func (h *AccountHandler) Bind(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var account models.Account
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&account).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "账号不存在")
return
}
if account.Status == "active" {
response.Fail(c, http.StatusConflict, 3002, "账号已绑定")
return
}
challenge := models.Challenge{
WorkspaceID: c.GetUint64("wsid"),
AccountID: account.ID,
Platform: account.Platform,
Kind: "pending",
Status: "active",
Prompt: "等待客户端发起扫码,稍后此处展示二维码",
ExpiresAt: time.Now().Add(30 * time.Minute),
}
err = h.DB.Transaction(func(tx *gorm.DB) error {
if err = tx.Create(&challenge).Error; err != nil {
return err
}
return tx.Model(&account).Update("status", "binding").Error
})
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "发起绑定失败")
return
}
// 定向下发:优先账号指定设备,否则工作空间首个在线设备
if h.Hub != nil {
target := account.AgentDeviceID
if target == 0 || !h.Hub.AgentOnline(target) {
if online := h.Hub.OnlineDevices(c.GetUint64("wsid")); len(online) > 0 {
target = online[0]
}
}
if target != 0 {
env := proto.NewEnvelope(uuid.NewString(), proto.TypeChallenge, proto.Challenge{
ChallengeID: strconv.FormatUint(challenge.ID, 10),
AccountID: strconv.FormatUint(account.ID, 10),
Platform: account.Platform,
Kind: "qr",
Prompt: "请使用手机客户端扫码登录 " + account.Platform,
ExpiresAt: challenge.ExpiresAt.UnixMilli(),
})
_ = h.Hub.SendToAgent(target, env)
}
}
response.Audit(c, "account.bind", "account:"+itoa(account.ID), gin.H{"challengeId": challenge.ID})
response.OK(c, gin.H{"challengeId": challenge.ID, "account": account.ID, "status": "binding"})
}
+143
View File
@@ -0,0 +1,143 @@
package handlers
import (
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
"everypublish/server/internal/ws"
)
// AdminHandler 平台管理处理器(AdminRequired 守卫)
type AdminHandler struct {
DB *gorm.DB
Hub *ws.Hub
}
type tenantRow struct {
models.Workspace
OwnerEmail string `json:"ownerEmail"`
MemberCount int64 `json:"memberCount"`
DeviceCount int64 `json:"deviceCount"`
}
// Tenants 租户列表(含 owner 邮箱与成员/设备数)
func (h *AdminHandler) Tenants(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.Workspace{})
if s := c.Query("status"); s != "" {
q = q.Where("status = ?", s)
}
var total int64
q.Count(&total)
var wsList []models.Workspace
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&wsList)
rows := make([]tenantRow, 0, len(wsList))
for _, ws := range wsList {
row := tenantRow{Workspace: ws}
var owner models.User
if err := h.DB.First(&owner, ws.OwnerID).Error; err == nil {
row.OwnerEmail = owner.Email
}
h.DB.Model(&models.Member{}).Where("workspace_id = ?", ws.ID).Count(&row.MemberCount)
h.DB.Model(&models.AgentDevice{}).Where("workspace_id = ?", ws.ID).Count(&row.DeviceCount)
rows = append(rows, row)
}
response.OK(c, gin.H{"list": rows, "total": total, "page": page, "size": size})
}
// TenantStatus 启用/停用租户
func (h *AdminHandler) TenantStatus(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req struct {
Status string `json:"status" binding:"required"`
}
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
if req.Status != "active" && req.Status != "suspended" {
response.Fail(c, http.StatusBadRequest, 1001, "状态不合法")
return
}
var ws models.Workspace
if err = h.DB.First(&ws, id).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "租户不存在")
return
}
if err = h.DB.Model(&ws).Update("status", req.Status).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "admin.tenant.status", "workspace:"+itoa(id), gin.H{"status": req.Status})
response.OK(c, nil)
}
type agentRow struct {
models.AgentDevice
WorkspaceName string `json:"workspaceName"`
}
// Agents 全局 Agent 总览
func (h *AdminHandler) Agents(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
var total int64
h.DB.Model(&models.AgentDevice{}).Count(&total)
var devices []models.AgentDevice
h.DB.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&devices)
rows := make([]agentRow, 0, len(devices))
for _, d := range devices {
row := agentRow{AgentDevice: d}
var ws models.Workspace
if err := h.DB.First(&ws, d.WorkspaceID).Error; err == nil {
row.WorkspaceName = ws.Name
}
rows = append(rows, row)
}
response.OK(c, gin.H{"list": rows, "total": total, "page": page, "size": size})
}
// AgentRevoke 全局强制吊销
func (h *AdminHandler) AgentRevoke(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var device models.AgentDevice
if err = h.DB.First(&device, id).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "设备不存在")
return
}
if err = h.DB.Model(&device).Updates(map[string]interface{}{"revoked": true, "status": "offline"}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "吊销失败")
return
}
if h.Hub != nil {
h.Hub.KickAgent(device.ID)
}
response.Audit(c, "admin.agent.revoke", "device:"+itoa(id), nil)
response.OK(c, nil)
}
+147
View File
@@ -0,0 +1,147 @@
package handlers
import (
"crypto/rand"
"errors"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
"everypublish/server/internal/ws"
)
// AgentHandler 设备配对与列表
type AgentHandler struct {
DB *gorm.DB
Hub *ws.Hub
}
const pairAlphabet = "ABCDEFGHJKMNPQRSTUVWXYZ23456789"
var errPairCodeUsed = errors.New("pair code already used")
func randCode(n int) string {
b := make([]byte, n)
_, _ = rand.Read(b)
for i := range b {
b[i] = pairAlphabet[int(b[i])%len(pairAlphabet)]
}
return string(b)
}
// PairCode 生成配对码(5 分钟一次性)
func (h *AgentHandler) PairCode(c *gin.Context) {
pc := models.PairingCode{
WorkspaceID: c.GetUint64("wsid"),
Code: randCode(6),
ExpiresAt: time.Now().Add(5 * time.Minute),
}
if err := h.DB.Create(&pc).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "生成配对码失败")
return
}
response.Audit(c, "agent.pair_code", "pairing:"+itoa(pc.ID), gin.H{"code": pc.Code})
response.OK(c, gin.H{"code": pc.Code, "expiresAt": pc.ExpiresAt})
}
// Pair 设备配对(无鉴权,凭一次性码;注册 Ed25519 公钥)
func (h *AgentHandler) Pair(c *gin.Context) {
var req struct {
Code string `json:"code" binding:"required"`
DeviceName string `json:"deviceName" binding:"required"`
OS string `json:"os"`
Version string `json:"version"`
PublicKey string `json:"publicKey" binding:"required"`
}
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var pc models.PairingCode
if err := h.DB.Where("code = ? and used = ?", req.Code, false).First(&pc).Error; err != nil {
response.Fail(c, http.StatusBadRequest, 2006, "配对码无效")
return
}
if time.Now().After(pc.ExpiresAt) {
response.Fail(c, http.StatusGone, 2006, "配对码已过期")
return
}
device := models.AgentDevice{
WorkspaceID: pc.WorkspaceID,
Name: req.DeviceName,
OS: req.OS,
Version: req.Version,
PublicKey: req.PublicKey,
Status: "offline",
PairedAt: time.Now(),
}
err := h.DB.Transaction(func(tx *gorm.DB) error {
// 原子占用:仅当仍 unused 且未过期时置 used=true;RowsAffected==1 才视为抢到,
// 防止并发用同一码配出多台设备。
res := tx.Model(&models.PairingCode{}).
Where("id = ? and used = ?", pc.ID, false).
Update("used", true)
if res.Error != nil {
return res.Error
}
if res.RowsAffected != 1 {
return errPairCodeUsed
}
if err := tx.Create(&device).Error; err != nil {
return err
}
if err := tx.Model(&models.PairingCode{}).Where("id = ?", pc.ID).
Update("device_id", device.ID).Error; err != nil {
return err
}
return nil
})
if err != nil {
if err == errPairCodeUsed {
response.Fail(c, http.StatusBadRequest, 2006, "配对码已被使用")
return
}
response.Fail(c, http.StatusInternalServerError, 5000, "配对失败")
return
}
response.OK(c, gin.H{"deviceId": device.ID, "workspaceId": device.WorkspaceID})
}
// Devices 设备列表(在线状态来自 Hub,此处以 DB 状态兜底)
func (h *AgentHandler) Devices(c *gin.Context) {
var list []models.AgentDevice
h.DB.Where("workspace_id = ?", c.GetUint64("wsid")).Order("id desc").Find(&list)
response.OK(c, gin.H{"list": list})
}
// Revoke 吊销设备并强制断开(owner/admin)
func (h *AgentHandler) Revoke(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可吊销设备")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var device models.AgentDevice
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&device).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "设备不存在")
return
}
if err = h.DB.Model(&device).Updates(map[string]interface{}{"revoked": true, "status": "offline"}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "吊销失败")
return
}
if h.Hub != nil {
h.Hub.KickAgent(device.ID)
}
response.Audit(c, "agent.revoke", "device:"+itoa(id), nil)
response.OK(c, nil)
}
+40
View File
@@ -0,0 +1,40 @@
package handlers
import (
"strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// AuditHandler 审计日志处理器
type AuditHandler struct {
DB *gorm.DB
}
// List 审计日志(分页 + 筛选)
func (h *AuditHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.AuditLog{}).Where("workspace_id = ?", c.GetUint64("wsid"))
if uid := c.Query("userId"); uid != "" {
q = q.Where("user_id = ?", uid)
}
if action := c.Query("action"); action != "" {
q = q.Where("action = ?", action)
}
var total int64
q.Count(&total)
var list []models.AuditLog
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&list)
response.OK(c, gin.H{"list": list, "total": total, "page": page, "size": size})
}
+235
View File
@@ -0,0 +1,235 @@
package handlers
import (
"net/http"
"strings"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/auth"
"everypublish/server/internal/cache"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
)
// AuthHandler 认证处理器
type AuthHandler struct {
DB *gorm.DB
Rds *cache.Redis
Cfg *config.Config
}
type registerReq struct {
Email string `json:"email" binding:"required"`
Password string `json:"password" binding:"required,min=8"`
Nickname string `json:"nickname"`
}
type loginReq struct {
Email string `json:"email" binding:"required"`
Password string `json:"password" binding:"required"`
TotpCode string `json:"totpCode"`
}
type refreshReq struct {
RefreshToken string `json:"refreshToken" binding:"required"`
}
type logoutReq struct {
RefreshToken string `json:"refreshToken"`
}
// Register 注册:创建用户 + 默认工作空间 + owner 成员
func (h *AuthHandler) Register(c *gin.Context) {
var req registerReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法:"+err.Error())
return
}
req.Email = strings.ToLower(strings.TrimSpace(req.Email))
if !strings.Contains(req.Email, "@") {
response.Fail(c, http.StatusBadRequest, 1001, "邮箱格式不正确")
return
}
if req.Nickname == "" {
req.Nickname = strings.Split(req.Email, "@")[0]
}
var count int64
h.DB.Model(&models.User{}).Where("email = ?", req.Email).Count(&count)
if count > 0 {
response.Fail(c, http.StatusConflict, 2001, "邮箱已注册")
return
}
hash, err := auth.HashPassword(req.Password)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "服务内部错误")
return
}
err = h.DB.Transaction(func(tx *gorm.DB) error {
user := models.User{Email: req.Email, PasswordHash: hash, Nickname: req.Nickname, Role: "user", Status: "active"}
if err := tx.Create(&user).Error; err != nil {
return err
}
ws := models.Workspace{Name: req.Nickname + "的工作空间", OwnerID: user.ID, Plan: "free", Status: "active"}
if err := tx.Create(&ws).Error; err != nil {
return err
}
member := models.Member{WorkspaceID: ws.ID, UserID: user.ID, Role: "owner"}
if err := tx.Create(&member).Error; err != nil {
return err
}
c.Set("regUID", user.ID)
c.Set("regWSID", ws.ID)
return nil
})
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "注册失败")
return
}
uid, wsid := c.GetUint64("regUID"), c.GetUint64("regWSID")
access, refresh, err := h.issuePair(uid, wsid, "user", "owner")
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "签发令牌失败")
return
}
var user models.User
h.DB.First(&user, uid)
response.OK(c, gin.H{
"accessToken": access,
"refreshToken": refresh,
"user": userDTO(user),
"workspace": gin.H{"id": wsid, "name": req.Nickname + "的工作空间"},
})
}
// Login 登录(2FA 默认关;开启后需 totpCode)
func (h *AuthHandler) Login(c *gin.Context) {
var req loginReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var user models.User
if err := h.DB.Where("email = ?", strings.ToLower(strings.TrimSpace(req.Email))).First(&user).Error; err != nil {
response.Fail(c, http.StatusUnauthorized, 2002, "邮箱或密码错误")
return
}
ok, err := auth.VerifyPassword(req.Password, user.PasswordHash)
if err != nil || !ok {
response.Fail(c, http.StatusUnauthorized, 2002, "邮箱或密码错误")
return
}
if user.TOTPEnabled {
// 二期接入 TOTP 校验;一期默认关闭
response.Fail(c, http.StatusForbidden, 1006, "需要两步验证码")
return
}
if user.Status != "active" {
response.Fail(c, http.StatusForbidden, 2003, "账号已停用")
return
}
var member models.Member
if err := h.DB.Where("user_id = ?", user.ID).Order("id asc").First(&member).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "工作空间数据缺失")
return
}
access, refresh, err := h.issuePair(user.ID, member.WorkspaceID, user.Role, member.Role)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "签发令牌失败")
return
}
c.Set("wsid", member.WorkspaceID)
response.Audit(c, "auth.login", "user:"+itoa(user.ID), gin.H{"email": user.Email})
response.OK(c, gin.H{
"accessToken": access,
"refreshToken": refresh,
"user": userDTO(user),
"workspace": gin.H{"id": member.WorkspaceID},
})
}
// Refresh 刷新令牌轮换:旧 jti 吊销,签发新对
func (h *AuthHandler) Refresh(c *gin.Context) {
var req refreshReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
claims, err := auth.ParseRefresh(h.Cfg.JWTSecret, req.RefreshToken)
if err != nil {
response.Fail(c, http.StatusUnauthorized, 2003, "刷新令牌无效或已过期")
return
}
valid, err := h.Rds.ExistsRefresh(claims.UID, claims.JTI)
if err != nil || !valid {
response.Fail(c, http.StatusUnauthorized, 2003, "刷新令牌已失效")
return
}
if err = h.Rds.DeleteRefresh(claims.UID, claims.JTI); err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "服务内部错误")
return
}
var user models.User
if err = h.DB.First(&user, claims.UID).Error; err != nil {
response.Fail(c, http.StatusUnauthorized, 2003, "用户不存在")
return
}
var member models.Member
if err = h.DB.Where("user_id = ?", user.ID).Order("id asc").First(&member).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "工作空间数据缺失")
return
}
access, refresh, err := h.issuePair(user.ID, member.WorkspaceID, user.Role, member.Role)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "签发令牌失败")
return
}
response.OK(c, gin.H{"accessToken": access, "refreshToken": refresh})
}
// Logout 登出:吊销刷新令牌
func (h *AuthHandler) Logout(c *gin.Context) {
var req logoutReq
_ = c.ShouldBindJSON(&req)
if req.RefreshToken != "" {
if claims, err := auth.ParseRefresh(h.Cfg.JWTSecret, req.RefreshToken); err == nil {
_ = h.Rds.DeleteRefresh(claims.UID, claims.JTI)
}
}
response.Audit(c, "auth.logout", "user:"+itoa(c.GetUint64("uid")), nil)
response.OK(c, nil)
}
// Me 当前用户信息
func (h *AuthHandler) Me(c *gin.Context) {
var user models.User
if err := h.DB.First(&user, c.GetUint64("uid")).Error; err != nil {
response.Fail(c, http.StatusUnauthorized, 1002, "用户不存在")
return
}
response.OK(c, gin.H{"user": userDTO(user), "workspaceId": c.GetUint64("wsid"), "memberRole": c.GetString("mrole")})
}
// issuePair 签发令牌对并把 refresh jti 登记到 Redis
func (h *AuthHandler) issuePair(uid, wsid uint64, role, mrole string) (string, string, error) {
access, err := auth.IssueAccess(h.Cfg.JWTSecret, uid, wsid, role, mrole)
if err != nil {
return "", "", err
}
jti := uuid.NewString()
refresh, err := auth.IssueRefresh(h.Cfg.JWTSecret, uid, jti)
if err != nil {
return "", "", err
}
if err = h.Rds.SaveRefresh(uid, jti, auth.RefreshTTL); err != nil {
return "", "", err
}
return access, refresh, nil
}
func userDTO(u models.User) gin.H {
return gin.H{"id": u.ID, "email": u.Email, "nickname": u.Nickname, "role": u.Role, "totpEnabled": u.TOTPEnabled}
}
+114
View File
@@ -0,0 +1,114 @@
package handlers
import (
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// ChallengeHandler 二次验证挑战处理器
type ChallengeHandler struct {
DB *gorm.DB
}
// List 挑战列表(过期软置 expired)
func (h *ChallengeHandler) List(c *gin.Context) {
h.DB.Model(&models.Challenge{}).
Where("workspace_id = ? and status = ? and expires_at < ?", c.GetUint64("wsid"), "active", time.Now()).
Update("status", "expired")
q := h.DB.Model(&models.Challenge{}).Where("workspace_id = ?", c.GetUint64("wsid"))
if s := c.Query("status"); s != "" {
q = q.Where("status = ?", s)
}
if a := c.Query("accountId"); a != "" {
q = q.Where("account_id = ?", a)
}
var list []models.Challenge
q.Order("id desc").Limit(100).Find(&list)
response.OK(c, gin.H{"list": list})
}
// Solve 人工完成挑战(验证码/APP确认;扫码由 Agent 完成)
func (h *ChallengeHandler) Solve(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req struct {
Value string `json:"value"`
}
_ = c.ShouldBindJSON(&req)
var challenge models.Challenge
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&challenge).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "挑战不存在")
return
}
if challenge.Status != "active" {
response.Fail(c, http.StatusConflict, 3003, "挑战已结束")
return
}
now := time.Now()
if err = h.DB.Model(&challenge).Updates(map[string]interface{}{"status": "solved", "solved_at": &now}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "challenge.solve", "challenge:"+itoa(id), gin.H{"kind": challenge.Kind})
response.OK(c, nil)
}
// Resend 一键重发(重置过期时间,实际重发出 Agent 执行)
func (h *ChallengeHandler) Resend(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var challenge models.Challenge
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&challenge).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "挑战不存在")
return
}
if challenge.Status != "expired" && challenge.Status != "suspended" {
response.Fail(c, http.StatusConflict, 3003, "当前状态不可重发")
return
}
if err = h.DB.Model(&challenge).Updates(map[string]interface{}{
"status": "active", "expires_at": time.Now().Add(30 * time.Minute),
}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "challenge.resend", "challenge:"+itoa(id), nil)
response.OK(c, nil)
}
// Suspend 挂起(超时/人工)
func (h *ChallengeHandler) Suspend(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var challenge models.Challenge
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&challenge).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "挑战不存在")
return
}
if challenge.Status != "active" {
response.Fail(c, http.StatusConflict, 3003, "挑战已结束")
return
}
if err = h.DB.Model(&challenge).Update("status", "suspended").Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "challenge.suspend", "challenge:"+itoa(id), nil)
response.OK(c, nil)
}
+201
View File
@@ -0,0 +1,201 @@
package handlers
import (
"crypto/sha256"
"encoding/hex"
"io"
"net/http"
"os"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
)
// MaterialHandler 素材处理器(服务器本地盘一期;StorageDriver 二期 OSS)
type MaterialHandler struct {
DB *gorm.DB
Cfg *config.Config
}
// List 素材列表(分页 + 筛选)
func (h *MaterialHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.Material{}).Where("workspace_id = ? and status = ?", c.GetUint64("wsid"), "ready")
if k := c.Query("kind"); k != "" {
q = q.Where("kind = ?", k)
}
if g := c.Query("group"); g != "" {
q = q.Where("`group` = ?", g)
}
var total int64
q.Count(&total)
var list []models.Material
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&list)
response.OK(c, gin.H{"list": list, "total": total, "page": page, "size": size})
}
// Upload multipart 上传:流式计算 sha256,同工作区按 sha256 去重
func (h *MaterialHandler) Upload(c *gin.Context) {
if h.Cfg.MaxUploadBytes > 0 {
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, h.Cfg.MaxUploadBytes)
}
file, header, err := c.Request.FormFile("file")
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "文件过大或缺少文件字段 file")
return
}
defer file.Close()
wsid := c.GetUint64("wsid")
dir := filepath.Join(h.Cfg.StorageDir, itoa(wsid))
if err = os.MkdirAll(dir, 0o755); err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "存储初始化失败")
return
}
ext := strings.ToLower(filepath.Ext(header.Filename))
if ext == "" || len(ext) > 16 {
ext = ".bin"
}
key := filepath.Join(itoa(wsid), uuid.NewString()+ext)
abspath := filepath.Join(h.Cfg.StorageDir, key)
dst, err := os.Create(abspath)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "写盘失败")
return
}
hasher := sha256.New()
size, err := io.Copy(io.MultiWriter(dst, hasher), file)
dst.Close()
if err != nil {
_ = os.Remove(abspath)
response.Fail(c, http.StatusInternalServerError, 5000, "写入失败")
return
}
sum := hex.EncodeToString(hasher.Sum(nil))
// 去重:同工作区同 sha256 直接复用
var existing models.Material
if err = h.DB.Where("workspace_id = ? and sha256 = ? and status = ?", wsid, sum, "ready").First(&existing).Error; err == nil {
_ = os.Remove(abspath)
response.OK(c, gin.H{"material": existing, "dedup": true})
return
}
kind := "image"
if strings.HasPrefix(header.Header.Get("Content-Type"), "video/") {
kind = "video"
}
group := strings.TrimSpace(c.PostForm("group"))
tags := strings.TrimSpace(c.PostForm("tags"))
if tags == "" {
tags = "[]"
}
material := models.Material{
WorkspaceID: wsid,
Name: header.Filename,
Kind: kind,
Size: size,
SHA256: sum,
Mime: header.Header.Get("Content-Type"),
StorageKey: key,
Group: group,
Tags: tags,
Status: "ready",
CreatedBy: c.GetUint64("uid"),
}
if err = h.DB.Create(&material).Error; err != nil {
_ = os.Remove(abspath)
response.Fail(c, http.StatusInternalServerError, 5000, "入库失败")
return
}
response.Audit(c, "material.upload", "material:"+itoa(material.ID), gin.H{"name": header.Filename, "size": size})
response.OK(c, gin.H{"material": material, "dedup": false})
}
// URL 生成签名直链(10 分钟一次性)
func (h *MaterialHandler) URL(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var material models.Material
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&material).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "素材不存在")
return
}
token := models.TransferToken{
MaterialID: material.ID,
Token: uuid.NewString(),
ExpiresAt: time.Now().Add(10 * time.Minute),
}
if err = h.DB.Create(&token).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "生成链接失败")
return
}
response.OK(c, gin.H{"url": h.Cfg.BaseURL + "/api/v1/files/" + token.Token, "expiresAt": token.ExpiresAt})
}
// ServeFile 凭 token 取文件(一次性,10 分钟有效)
func (h *MaterialHandler) ServeFile(c *gin.Context) {
tokenStr := c.Param("token")
var token models.TransferToken
if err := h.DB.Where("token = ?", tokenStr).First(&token).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "链接无效")
return
}
if time.Now().After(token.ExpiresAt) {
h.DB.Delete(&token)
response.Fail(c, http.StatusGone, 3004, "链接已过期")
return
}
var material models.Material
if err := h.DB.First(&material, token.MaterialID).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "素材不存在")
return
}
// 一次性:取用即吊销
_ = h.DB.Delete(&token).Error
abspath := filepath.Join(h.Cfg.StorageDir, material.StorageKey)
clean := filepath.Clean(abspath)
if !strings.HasPrefix(clean, filepath.Clean(h.Cfg.StorageDir)) {
response.Fail(c, http.StatusForbidden, 1003, "非法路径")
return
}
c.FileAttachment(abspath, material.Name)
}
// Delete 删除素材(含文件)
func (h *MaterialHandler) Delete(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var material models.Material
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&material).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "素材不存在")
return
}
abspath := filepath.Join(h.Cfg.StorageDir, material.StorageKey)
_ = os.Remove(abspath)
if err = h.DB.Delete(&material).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "删除失败")
return
}
response.Audit(c, "material.delete", "material:"+itoa(id), nil)
response.OK(c, nil)
}
+190
View File
@@ -0,0 +1,190 @@
package handlers
import (
"net/http"
"strconv"
"strings"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/auth"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
)
// MemberHandler 成员处理器
type MemberHandler struct {
DB *gorm.DB
Cfg *config.Config
}
type inviteReq struct {
Email string `json:"email" binding:"required"`
Role string `json:"role" binding:"required"`
}
type joinReq struct {
Token string `json:"token" binding:"required"`
}
type roleReq struct {
Role string `json:"role" binding:"required"`
}
var memberRoles = map[string]bool{"admin": true, "operator": true, "reviewer": true, "viewer": true}
func isAdminRole(mrole string) bool { return mrole == "owner" || mrole == "admin" }
type memberRow struct {
models.Member
Email string `json:"email"`
Nickname string `json:"nickname"`
}
// List 成员列表(含用户资料)
func (h *MemberHandler) List(c *gin.Context) {
wsid := c.GetUint64("wsid")
var rows []memberRow
h.DB.Table("members").Select("members.*, users.email as email, users.nickname as nickname").
Joins("join users on users.id = members.user_id").
Where("members.workspace_id = ?", wsid).Order("members.id asc").Scan(&rows)
response.OK(c, gin.H{"list": rows})
}
// Invite 生成邀请链接(owner/admin)
func (h *MemberHandler) Invite(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可邀请成员")
return
}
var req inviteReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
req.Role = strings.ToLower(strings.TrimSpace(req.Role))
if !memberRoles[req.Role] {
response.Fail(c, http.StatusBadRequest, 1001, "角色不合法")
return
}
wsid := c.GetUint64("wsid")
var existing memberRow
err := h.DB.Table("members").Select("members.id").
Joins("join users on users.id = members.user_id").
Where("members.workspace_id = ? and users.email = ?", wsid, strings.ToLower(req.Email)).
Scan(&existing).Error
if err == nil && existing.ID > 0 {
response.Fail(c, http.StatusConflict, 2004, "该邮箱已是成员")
return
}
token, err := auth.IssueInvite(h.Cfg.JWTSecret, wsid, strings.ToLower(req.Email), req.Role)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "生成邀请失败")
return
}
response.Audit(c, "member.invite", "workspace:"+itoa(wsid), gin.H{"email": req.Email, "role": req.Role})
response.OK(c, gin.H{"token": token, "link": h.Cfg.PublicBaseURL + "/invite?t=" + token})
}
// Join 接受邀请(需登录,邮箱须与邀请一致)
func (h *MemberHandler) Join(c *gin.Context) {
var req joinReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
claims, err := auth.ParseInvite(h.Cfg.JWTSecret, req.Token)
if err != nil {
response.Fail(c, http.StatusBadRequest, 2005, "邀请链接无效或已过期")
return
}
var user models.User
if err = h.DB.First(&user, c.GetUint64("uid")).Error; err != nil {
response.Fail(c, http.StatusUnauthorized, 1002, "用户不存在")
return
}
if !strings.EqualFold(user.Email, claims.Email) {
response.Fail(c, http.StatusForbidden, 1003, "邀请链接与当前账号邮箱不匹配")
return
}
member := models.Member{WorkspaceID: claims.WorkspaceID, UserID: user.ID, Role: claims.Role}
err = h.DB.Where("workspace_id = ? and user_id = ?", claims.WorkspaceID, user.ID).First(&models.Member{}).Error
if err == nil {
response.Fail(c, http.StatusConflict, 2004, "你已是该工作空间成员")
return
}
if err = h.DB.Create(&member).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "加入失败")
return
}
response.Audit(c, "member.join", "workspace:"+itoa(claims.WorkspaceID), gin.H{"role": claims.Role})
response.OK(c, nil)
}
// UpdateRole 修改成员角色(owner/admin;不可改 owner)
func (h *MemberHandler) UpdateRole(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可修改角色")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req roleReq
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
req.Role = strings.ToLower(strings.TrimSpace(req.Role))
if !memberRoles[req.Role] && req.Role != "owner" {
response.Fail(c, http.StatusBadRequest, 1001, "角色不合法")
return
}
var member models.Member
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&member).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "成员不存在")
return
}
if member.Role == "owner" {
response.Fail(c, http.StatusForbidden, 1003, "不可修改 owner 角色")
return
}
if err = h.DB.Model(&member).Update("role", req.Role).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "member.role", "member:"+itoa(id), gin.H{"role": req.Role})
response.OK(c, nil)
}
// Remove 移除成员(owner/admin;不可移除 owner)
func (h *MemberHandler) Remove(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可移除成员")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var member models.Member
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&member).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "成员不存在")
return
}
if member.Role == "owner" {
response.Fail(c, http.StatusForbidden, 1003, "不可移除 owner")
return
}
if err = h.DB.Delete(&member).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "移除失败")
return
}
response.Audit(c, "member.remove", "member:"+itoa(id), nil)
response.OK(c, nil)
}
@@ -0,0 +1,57 @@
package handlers
import (
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// NotificationHandler 站内通知处理器
type NotificationHandler struct {
DB *gorm.DB
}
// List 我的通知 + 未读数
func (h *NotificationHandler) List(c *gin.Context) {
q := h.DB.Model(&models.Notification{}).Where("workspace_id = ? and user_id = ?", c.GetUint64("wsid"), c.GetUint64("uid"))
if k := c.Query("kind"); k != "" {
q = q.Where("kind = ?", k)
}
var unread int64
q.Where("`read` = ?", false).Count(&unread)
var list []models.Notification
q.Order("id desc").Limit(100).Find(&list)
response.OK(c, gin.H{"list": list, "unread": unread})
}
// Read 标记已读
func (h *NotificationHandler) Read(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
if err = h.DB.Model(&models.Notification{}).
Where("id = ? and user_id = ?", id, c.GetUint64("uid")).
Update("read", true).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.OK(c, nil)
}
// ReadAll 全部已读
func (h *NotificationHandler) ReadAll(c *gin.Context) {
if err := h.DB.Model(&models.Notification{}).
Where("workspace_id = ? and user_id = ?", c.GetUint64("wsid"), c.GetUint64("uid")).
Update("read", true).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.OK(c, nil)
}
+25
View File
@@ -0,0 +1,25 @@
package handlers
import (
"gorm.io/gorm"
"everypublish/server/internal/models"
)
// Notify 写入站内通知(批量用户)
func Notify(db *gorm.DB, wsid uint64, userIDs []uint64, kind, title, content string) error {
if len(userIDs) == 0 {
return nil
}
rows := make([]models.Notification, 0, len(userIDs))
for _, uid := range userIDs {
rows = append(rows, models.Notification{
WorkspaceID: wsid,
UserID: uid,
Kind: kind,
Title: title,
Content: content,
})
}
return db.Create(&rows).Error
}
+279
View File
@@ -0,0 +1,279 @@
package handlers
import (
"encoding/json"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
"everypublish/server/internal/task"
"everypublish/server/internal/ws"
)
// TaskHandler 发布任务处理器(状态机见 internal/task)
type TaskHandler struct {
DB *gorm.DB
Dispatcher *ws.Dispatcher
}
type taskReq struct {
Title string `json:"title" binding:"required"`
Content string `json:"content"`
Tags []string `json:"tags"`
AccountIDs []uint64 `json:"accountIds"`
MaterialIDs []uint64 `json:"materialIds"`
ScheduleAt *int64 `json:"scheduleAt"`
Priority int `json:"priority"`
}
func marshalJSONArray[T any](arr []T) string {
raw, _ := json.Marshal(arr)
if raw == nil {
return "[]"
}
return string(raw)
}
// List 任务列表(状态筛选 + 分页)
func (h *TaskHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.Task{}).Where("workspace_id = ?", c.GetUint64("wsid"))
if s := c.Query("status"); s != "" {
q = q.Where("status = ?", s)
}
if s := c.Query("schedule"); s != "" {
if s == "scheduled" {
q = q.Where("schedule_at is not null and schedule_at > ?", time.Now())
} else if s == "immediate" {
q = q.Where("schedule_at is null or schedule_at <= ?", time.Now())
}
}
var total int64
q.Count(&total)
var list []models.Task
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&list)
response.OK(c, gin.H{"list": list, "total": total, "page": page, "size": size})
}
// Create 新建任务(draft)
func (h *TaskHandler) Create(c *gin.Context) {
var req taskReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法:"+err.Error())
return
}
if len(req.AccountIDs) == 0 {
response.Fail(c, http.StatusBadRequest, 1001, "至少选择一个平台账号")
return
}
if len(req.MaterialIDs) == 0 {
response.Fail(c, http.StatusBadRequest, 1001, "至少选择一个素材")
return
}
priority := req.Priority
if priority < 1 || priority > 10 {
priority = 5
}
row := models.Task{
WorkspaceID: c.GetUint64("wsid"),
Title: req.Title,
Content: req.Content,
Tags: marshalJSONArray(req.Tags),
AccountIDs: marshalJSONArray(req.AccountIDs),
MaterialIDs: marshalJSONArray(req.MaterialIDs),
Status: task.Draft,
Priority: priority,
MaxRetry: 3,
CreatedBy: c.GetUint64("uid"),
}
if req.ScheduleAt != nil {
ts := time.UnixMilli(*req.ScheduleAt)
row.ScheduleAt = &ts
}
if err := h.DB.Create(&row).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "创建失败")
return
}
response.Audit(c, "task.create", "task:"+itoa(row.ID), gin.H{"title": req.Title})
response.OK(c, row)
}
// Get 任务详情
func (h *TaskHandler) Get(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var row models.Task
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&row).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "任务不存在")
return
}
response.OK(c, row)
}
// Update 编辑草稿(仅 draft/rejected)
func (h *TaskHandler) Update(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req taskReq
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var row models.Task
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&row).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "任务不存在")
return
}
if row.Status != task.Draft && row.Status != task.Rejected {
response.Fail(c, http.StatusConflict, 3005, "仅草稿或已驳回任务可编辑")
return
}
updates := map[string]interface{}{
"title": req.Title,
"content": req.Content,
"tags": marshalJSONArray(req.Tags),
"account_ids": marshalJSONArray(req.AccountIDs),
"material_ids": marshalJSONArray(req.MaterialIDs),
}
if req.ScheduleAt != nil {
ts := time.UnixMilli(*req.ScheduleAt)
updates["schedule_at"] = &ts
}
if err = h.DB.Model(&row).Updates(updates).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "task.update", "task:"+itoa(id), gin.H{"title": req.Title})
response.OK(c, nil)
}
// Delete 删除草稿
func (h *TaskHandler) Delete(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var row models.Task
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&row).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "任务不存在")
return
}
if row.Status != task.Draft {
response.Fail(c, http.StatusConflict, 3005, "仅草稿可删除")
return
}
if err = h.DB.Delete(&row).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "删除失败")
return
}
response.Audit(c, "task.delete", "task:"+itoa(id), nil)
response.OK(c, nil)
}
// apply 状态转移并落库
func (h *TaskHandler) apply(c *gin.Context, action task.Action, extra map[string]interface{}) *models.Task {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return nil
}
var row models.Task
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&row).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "任务不存在")
return nil
}
next, ok := task.Next(row.Status, action)
if !ok {
response.Fail(c, http.StatusConflict, 3005, "当前状态("+row.Status+")不允许该操作("+string(action)+")")
return nil
}
updates := map[string]interface{}{"status": next}
for k, v := range extra {
updates[k] = v
}
if err = h.DB.Model(&row).Updates(updates).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return nil
}
row.Status = next
return &row
}
// Submit 提交审核
func (h *TaskHandler) Submit(c *gin.Context) {
if row := h.apply(c, task.ActionSubmit, nil); row != nil {
response.Audit(c, "task.submit", "task:"+itoa(row.ID), nil)
response.OK(c, row)
}
}
// Approve 审核通过 → queued
func (h *TaskHandler) Approve(c *gin.Context) {
if row := h.apply(c, task.ActionApprove, nil); row != nil {
_ = Notify(h.DB, c.GetUint64("wsid"), []uint64{row.CreatedBy}, "task", "任务已通过审核", row.Title)
response.Audit(c, "task.approve", "task:"+itoa(row.ID), nil)
if h.Dispatcher != nil {
_ = h.Dispatcher.TryDispatch(row.ID) // 在线设备立即下发
}
response.OK(c, row)
}
}
// Reject 驳回(附批注)
func (h *TaskHandler) Reject(c *gin.Context) {
var req struct {
Note string `json:"note"`
}
_ = c.ShouldBindJSON(&req)
if row := h.apply(c, task.ActionReject, map[string]interface{}{"review_note": req.Note, "reviewed_by": c.GetUint64("uid")}); row != nil {
_ = Notify(h.DB, c.GetUint64("wsid"), []uint64{row.CreatedBy}, "task", "任务被驳回", req.Note)
response.Audit(c, "task.reject", "task:"+itoa(row.ID), gin.H{"note": req.Note})
response.OK(c, row)
}
}
// Resubmit 驳回后重新提交
func (h *TaskHandler) Resubmit(c *gin.Context) {
if row := h.apply(c, task.ActionResubmit, nil); row != nil {
response.Audit(c, "task.resubmit", "task:"+itoa(row.ID), nil)
response.OK(c, row)
}
}
// Retry 失败重试
func (h *TaskHandler) Retry(c *gin.Context) {
if row := h.apply(c, task.ActionRetry, map[string]interface{}{"error_message": ""}); row != nil {
response.Audit(c, "task.retry", "task:"+itoa(row.ID), nil)
if h.Dispatcher != nil {
_ = h.Dispatcher.TryDispatch(row.ID)
}
response.OK(c, row)
}
}
// Cancel 取消
func (h *TaskHandler) Cancel(c *gin.Context) {
if row := h.apply(c, task.ActionCancel, nil); row != nil {
response.Audit(c, "task.cancel", "task:"+itoa(row.ID), nil)
response.OK(c, row)
}
}
+7
View File
@@ -0,0 +1,7 @@
package handlers
import "strconv"
func itoa(v uint64) string {
return strconv.FormatUint(v, 10)
}
+84
View File
@@ -0,0 +1,84 @@
package handlers
import (
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// WorkspaceHandler 工作空间处理器
type WorkspaceHandler struct {
DB *gorm.DB
}
type workspaceReq struct {
Name string `json:"name" binding:"required"`
}
// Create 创建工作空间(创建者为 owner)
func (h *WorkspaceHandler) Create(c *gin.Context) {
var req workspaceReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var ws models.Workspace
err := h.DB.Transaction(func(tx *gorm.DB) error {
ws = models.Workspace{Name: req.Name, OwnerID: c.GetUint64("uid"), Plan: "free", Status: "active"}
if err := tx.Create(&ws).Error; err != nil {
return err
}
member := models.Member{WorkspaceID: ws.ID, UserID: c.GetUint64("uid"), Role: "owner"}
return tx.Create(&member).Error
})
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "创建失败")
return
}
response.Audit(c, "workspace.create", "workspace:"+itoa(ws.ID), gin.H{"name": req.Name})
response.OK(c, ws)
}
// List 我所在的工作空间
func (h *WorkspaceHandler) List(c *gin.Context) {
var members []models.Member
h.DB.Where("user_id = ?", c.GetUint64("uid")).Find(&members)
ids := make([]uint64, 0, len(members))
for _, m := range members {
ids = append(ids, m.WorkspaceID)
}
var list []models.Workspace
if len(ids) > 0 {
h.DB.Where("id in ?", ids).Order("id asc").Find(&list)
}
response.OK(c, gin.H{"list": list})
}
// Update 修改工作空间资料(owner/admin)
func (h *WorkspaceHandler) Update(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可修改工作空间")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req workspaceReq
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
if err = h.DB.Model(&models.Workspace{}).Where("id = ?", id).Update("name", req.Name).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "workspace.update", "workspace:"+itoa(id), gin.H{"name": req.Name})
response.OK(c, nil)
}
+38
View File
@@ -0,0 +1,38 @@
package middleware
import (
"encoding/json"
"net/http"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// Audit 审计埋点:处理器调用 response.Audit(c, action, resource, detail) 声明,
// 本中间件在请求成功(2xx)后落库。
func Audit(db *gorm.DB) gin.HandlerFunc {
return func(c *gin.Context) {
c.Next()
metaRaw, exists := c.Get("audit")
if !exists {
return
}
meta, ok := metaRaw.(response.AuditMeta)
if !ok || c.Writer.Status() >= http.StatusBadRequest {
return
}
detail, _ := json.Marshal(meta.Detail)
log := models.AuditLog{
WorkspaceID: c.GetUint64("wsid"),
UserID: c.GetUint64("uid"),
Action: meta.Action,
Resource: meta.Resource,
Detail: string(detail),
IP: c.ClientIP(),
}
_ = db.Create(&log).Error
}
}
+47
View File
@@ -0,0 +1,47 @@
package middleware
import (
"net/http"
"strings"
"github.com/gin-gonic/gin"
"everypublish/server/internal/api/response"
"everypublish/server/internal/auth"
"everypublish/server/internal/config"
)
// AuthRequired 校验访问令牌并把声明写入上下文
func AuthRequired(cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
header := c.GetHeader("Authorization")
if !strings.HasPrefix(header, "Bearer ") {
response.Fail(c, http.StatusUnauthorized, 1002, "未登录或令牌缺失")
c.Abort()
return
}
claims, err := auth.ParseAccess(cfg.JWTSecret, strings.TrimPrefix(header, "Bearer "))
if err != nil {
response.Fail(c, http.StatusUnauthorized, 1002, "令牌无效或已过期")
c.Abort()
return
}
c.Set("uid", claims.UID)
c.Set("wsid", claims.WorkspaceID)
c.Set("role", claims.Role)
c.Set("mrole", claims.MemberRole)
c.Next()
}
}
// AdminRequired 平台管理员守卫(D7 挂载 admin 路由)
func AdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
if c.GetString("role") != "admin" {
response.Fail(c, http.StatusForbidden, 1003, "无平台管理权限")
c.Abort()
return
}
c.Next()
}
}
+29
View File
@@ -0,0 +1,29 @@
package response
import (
"net/http"
"github.com/gin-gonic/gin"
)
// OK 统一成功返回
func OK(c *gin.Context, data interface{}) {
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": data})
}
// Fail 统一失败返回
func Fail(c *gin.Context, httpCode int, bizCode int, msg string) {
c.JSON(httpCode, gin.H{"code": bizCode, "message": msg, "data": nil})
}
// AuditMeta 处理器埋点元数据
type AuditMeta struct {
Action string `json:"action"`
Resource string `json:"resource"`
Detail interface{} `json:"detail"`
}
// Audit 审计埋点声明(由 middleware.Audit 中间件落库)
func Audit(c *gin.Context, action, resource string, detail interface{}) {
c.Set("audit", AuditMeta{Action: action, Resource: resource, Detail: detail})
}
+156
View File
@@ -0,0 +1,156 @@
package api
import (
"crypto/ed25519"
"net/http"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/handlers"
"everypublish/server/internal/api/middleware"
"everypublish/server/internal/cache"
"everypublish/server/internal/config"
"everypublish/server/internal/ws"
)
// Router 装配 gin 路由
func Router(cfg *config.Config, db *gorm.DB, rds *cache.Redis, hub *ws.Hub, serverKey ed25519.PrivateKey, dsp *ws.Dispatcher) *gin.Engine {
r := gin.New()
r.Use(gin.Logger(), gin.Recovery())
r.GET("/health", func(c *gin.Context) {
dbOK := false
if sqlDB, err := db.DB(); err == nil {
dbOK = sqlDB.Ping() == nil
}
redisOK := rds.Ping() == nil
c.JSON(http.StatusOK, gin.H{
"status": "ok",
"service": "everypublish-server",
"time": time.Now().Format(time.RFC3339),
"db": dbOK,
"redis": redisOK,
})
})
// WSS 双通道:Agent 设备签名 / 浏览器 JWT
r.GET("/ws/agent", ws.HandleAgent(cfg, db, hub, serverKey))
r.GET("/ws/browser", ws.HandleBrowser(cfg, hub))
api := r.Group("/api/v1")
api.Use(middleware.Audit(db))
authH := &handlers.AuthHandler{DB: db, Rds: rds, Cfg: cfg}
wsH := &handlers.WorkspaceHandler{DB: db}
mbH := &handlers.MemberHandler{DB: db, Cfg: cfg}
adH := &handlers.AuditHandler{DB: db}
accH := &handlers.AccountHandler{DB: db, Hub: hub}
matH := &handlers.MaterialHandler{DB: db, Cfg: cfg}
tskH := &handlers.TaskHandler{DB: db, Dispatcher: dsp}
chlH := &handlers.ChallengeHandler{DB: db}
ntfH := &handlers.NotificationHandler{DB: db}
agH := &handlers.AgentHandler{DB: db, Hub: hub}
auth := api.Group("/auth")
{
auth.POST("/register", authH.Register)
auth.POST("/login", authH.Login)
auth.POST("/refresh", authH.Refresh)
auth.POST("/logout", middleware.AuthRequired(cfg), authH.Logout)
auth.GET("/me", middleware.AuthRequired(cfg), authH.Me)
}
wsg := api.Group("/workspaces", middleware.AuthRequired(cfg))
{
wsg.POST("", wsH.Create)
wsg.GET("", wsH.List)
wsg.PUT("/:id", wsH.Update)
}
mb := api.Group("/members", middleware.AuthRequired(cfg))
{
mb.GET("", mbH.List)
mb.POST("/invite", mbH.Invite)
mb.POST("/join", mbH.Join)
mb.PUT("/:id/role", mbH.UpdateRole)
mb.DELETE("/:id", mbH.Remove)
}
ad := api.Group("/audit-logs", middleware.AuthRequired(cfg))
{
ad.GET("", adH.List)
}
acc := api.Group("/accounts", middleware.AuthRequired(cfg))
{
acc.GET("", accH.List)
acc.POST("", accH.Create)
acc.PUT("/:id", accH.Update)
acc.DELETE("/:id", accH.Delete)
acc.POST("/:id/bind", accH.Bind)
}
mat := api.Group("/materials", middleware.AuthRequired(cfg))
{
mat.GET("", matH.List)
mat.POST("", matH.Upload)
mat.DELETE("/:id", matH.Delete)
mat.GET("/:id/url", matH.URL)
}
files := api.Group("/files")
{
files.GET("/:token", matH.ServeFile)
}
tsk := api.Group("/tasks", middleware.AuthRequired(cfg))
{
tsk.GET("", tskH.List)
tsk.POST("", tskH.Create)
tsk.GET("/:id", tskH.Get)
tsk.PUT("/:id", tskH.Update)
tsk.DELETE("/:id", tskH.Delete)
tsk.POST("/:id/submit", tskH.Submit)
tsk.POST("/:id/approve", tskH.Approve)
tsk.POST("/:id/reject", tskH.Reject)
tsk.POST("/:id/resubmit", tskH.Resubmit)
tsk.POST("/:id/retry", tskH.Retry)
tsk.POST("/:id/cancel", tskH.Cancel)
}
chl := api.Group("/challenges", middleware.AuthRequired(cfg))
{
chl.GET("", chlH.List)
chl.POST("/:id/solve", chlH.Solve)
chl.POST("/:id/resend", chlH.Resend)
chl.POST("/:id/suspend", chlH.Suspend)
}
ntf := api.Group("/notifications", middleware.AuthRequired(cfg))
{
ntf.GET("", ntfH.List)
ntf.POST("/:id/read", ntfH.Read)
ntf.POST("/read-all", ntfH.ReadAll)
}
ag := api.Group("/agent")
{
ag.POST("/pair-code", middleware.AuthRequired(cfg), agH.PairCode)
ag.GET("/devices", middleware.AuthRequired(cfg), agH.Devices)
ag.PUT("/devices/:id/revoke", middleware.AuthRequired(cfg), agH.Revoke)
}
api.POST("/agent/pair", agH.Pair)
admH := &handlers.AdminHandler{DB: db, Hub: hub}
adm := api.Group("/admin", middleware.AuthRequired(cfg), middleware.AdminRequired())
{
adm.GET("/tenants", admH.Tenants)
adm.PUT("/tenants/:id/status", admH.TenantStatus)
adm.GET("/agents", admH.Agents)
adm.PUT("/agents/:id/revoke", admH.AgentRevoke)
}
return r
}
+250
View File
@@ -0,0 +1,250 @@
package api_test
import (
"bytes"
"crypto/ed25519"
"crypto/rand"
"encoding/json"
"fmt"
"net/http/httptest"
"os"
"testing"
"github.com/gin-gonic/gin"
gormmysql "gorm.io/driver/mysql"
"gorm.io/gorm"
"everypublish/server/internal/api"
"everypublish/server/internal/cache"
"everypublish/server/internal/config"
"everypublish/server/internal/db"
"everypublish/server/internal/models"
"everypublish/server/internal/ws"
)
var (
testRouter *gin.Engine
testDB *gorm.DB
testHub *ws.Hub
)
func TestMain(m *testing.M) {
adminDSN := os.Getenv("MYSQL_ADMIN_DSN")
if adminDSN == "" {
adminDSN = "root:everypublish@tcp(127.0.0.1:3306)/?charset=utf8mb4&parseTime=True&loc=Local"
}
adm, err := gorm.Open(gormmysql.Open(adminDSN), &gorm.Config{})
if err != nil {
fmt.Println("skip: mysql admin connect failed:", err)
os.Exit(0)
}
if err = adm.Exec("CREATE DATABASE IF NOT EXISTS everypublish_test CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci").Error; err != nil {
fmt.Println("skip: create test db failed:", err)
os.Exit(0)
}
testDSN := "root:everypublish@tcp(127.0.0.1:3306)/everypublish_test?charset=utf8mb4&parseTime=True&loc=Local"
testDB, err = gorm.Open(gormmysql.Open(testDSN), &gorm.Config{})
if err != nil {
fmt.Println("skip: test db connect failed:", err)
os.Exit(0)
}
resetTestDB()
cfg := &config.Config{
ServerAddr: ":0",
MySQLDSN: testDSN,
RedisAddr: "127.0.0.1:6379",
JWTSecret: "test-secret",
BaseURL: "http://127.0.0.1:8090",
StorageDir: "/tmp/everypublish-test",
}
rds := cache.New(cfg)
testHub = ws.NewHub()
dsp := ws.NewDispatcher(testDB, testHub, cfg)
_, serverKey, _ := ed25519.GenerateKey(rand.Reader)
testRouter = api.Router(cfg, testDB, rds, testHub, serverKey, dsp)
code := m.Run()
os.Exit(code)
}
// resetTestDB 清空并重建表 + 清空连接中枢(保证 -count 多轮可重复)
func resetTestDB() {
if testHub != nil {
testHub.Clear()
}
for _, t := range []interface{}{
&models.TransferToken{}, &models.PairingCode{}, &models.Notification{}, &models.AuditLog{},
&models.Challenge{}, &models.AgentDevice{}, &models.Task{}, &models.Material{},
&models.Account{}, &models.Member{}, &models.Workspace{}, &models.User{},
} {
testDB.Migrator().DropTable(t)
}
_ = db.AutoMigrate(testDB)
}
type envelope struct {
Code int `json:"code"`
Message string `json:"message"`
Data json.RawMessage `json:"data"`
}
func doReq(t *testing.T, method, path, token string, body interface{}) (int, envelope) {
t.Helper()
var buf bytes.Buffer
if body != nil {
raw, _ := json.Marshal(body)
buf.Write(raw)
}
req := httptest.NewRequest(method, path, &buf)
req.Header.Set("Content-Type", "application/json")
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
w := httptest.NewRecorder()
testRouter.ServeHTTP(w, req)
var env envelope
_ = json.Unmarshal(w.Body.Bytes(), &env)
return w.Code, env
}
type tokens struct {
AccessToken string `json:"accessToken"`
RefreshToken string `json:"refreshToken"`
}
func register(t *testing.T, email string) (string, string) {
t.Helper()
code, env := doReq(t, "POST", "/api/v1/auth/register", "", gin.H{"email": email, "password": "password-123", "nickname": "测试用户"})
if code != 200 || env.Code != 0 {
t.Fatalf("register failed: %d %s", code, env.Message)
}
var tk tokens
_ = json.Unmarshal(env.Data, &tk)
return tk.AccessToken, tk.RefreshToken
}
func TestAuthFlow(t *testing.T) {
resetTestDB()
access, refresh := register(t, "alice@test.com")
// 登录
code, env := doReq(t, "POST", "/api/v1/auth/login", "", gin.H{"email": "alice@test.com", "password": "password-123"})
if code != 200 || env.Code != 0 {
t.Fatalf("login failed: %d %s", code, env.Message)
}
// me
code, env = doReq(t, "GET", "/api/v1/auth/me", access, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("me failed: %d %s", code, env.Message)
}
// 错误密码
code, env = doReq(t, "POST", "/api/v1/auth/login", "", gin.H{"email": "alice@test.com", "password": "wrong-pass-1"})
if code != 401 || env.Code != 2002 {
t.Fatalf("wrong password should 2002, got %d %d", code, env.Code)
}
// 刷新轮换
code, env = doReq(t, "POST", "/api/v1/auth/refresh", "", gin.H{"refreshToken": refresh})
if code != 200 || env.Code != 0 {
t.Fatalf("refresh failed: %d %s", code, env.Message)
}
var tk tokens
_ = json.Unmarshal(env.Data, &tk)
newAccess, newRefresh := tk.AccessToken, tk.RefreshToken
// 旧 refresh 已被吊销
code, env = doReq(t, "POST", "/api/v1/auth/refresh", "", gin.H{"refreshToken": refresh})
if code != 401 || env.Code != 2003 {
t.Fatalf("old refresh should be revoked, got %d %d", code, env.Code)
}
// 登出后 refresh 失效
code, _ = doReq(t, "POST", "/api/v1/auth/logout", newAccess, gin.H{"refreshToken": newRefresh})
if code != 200 {
t.Fatalf("logout failed: %d", code)
}
code, env = doReq(t, "POST", "/api/v1/auth/refresh", "", gin.H{"refreshToken": newRefresh})
if code != 401 || env.Code != 2003 {
t.Fatalf("refresh after logout should fail, got %d %d", code, env.Code)
}
}
func TestWorkspaceAndMemberFlow(t *testing.T) {
resetTestDB()
aAccess, _ := register(t, "boss@test.com")
bAccess, _ := register(t, "worker@test.com")
// A 建第二个工作空间
code, env := doReq(t, "POST", "/api/v1/workspaces", aAccess, gin.H{"name": "第二空间"})
if code != 200 || env.Code != 0 {
t.Fatalf("create ws failed: %d %s", code, env.Message)
}
// A 邀请 B(viewer)
code, env = doReq(t, "POST", "/api/v1/members/invite", aAccess, gin.H{"email": "worker@test.com", "role": "viewer"})
if code != 200 || env.Code != 0 {
t.Fatalf("invite failed: %d %s", code, env.Message)
}
var inv struct {
Token string `json:"token"`
Link string `json:"link"`
}
_ = json.Unmarshal(env.Data, &inv)
if inv.Token == "" || inv.Link == "" {
t.Fatal("invite token/link missing")
}
// B 用错误邮箱邀请应失败(换 A 邀请 C 不存在的邮箱,B 接受会失败)
code, env = doReq(t, "POST", "/api/v1/members/join", bAccess, gin.H{"token": inv.Token})
if code != 200 || env.Code != 0 {
t.Fatalf("join failed: %d %s", code, env.Message)
}
// A 查成员:2 人
code, env = doReq(t, "GET", "/api/v1/members", aAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("list members failed: %d %s", code, env.Message)
}
var ml struct {
List []models.Member `json:"list"`
}
_ = json.Unmarshal(env.Data, &ml)
if len(ml.List) != 2 {
t.Fatalf("expect 2 members, got %d", len(ml.List))
}
var bMember uint64
var ownerMember uint64
for _, m := range ml.List {
if m.Role == "viewer" {
bMember = m.ID
}
if m.Role == "owner" {
ownerMember = m.ID
}
}
// A 改 B 为 operator
code, env = doReq(t, "PUT", fmt.Sprintf("/api/v1/members/%d/role", bMember), aAccess, gin.H{"role": "operator"})
if code != 200 || env.Code != 0 {
t.Fatalf("update role failed: %d %s", code, env.Message)
}
// owner 不可被移除
code, _ = doReq(t, "DELETE", fmt.Sprintf("/api/v1/members/%d", ownerMember), aAccess, nil)
if code != 403 {
t.Fatalf("remove owner should 403, got %d", code)
}
// owner 角色不可被修改
code, _ = doReq(t, "PUT", fmt.Sprintf("/api/v1/members/%d/role", ownerMember), aAccess, gin.H{"role": "viewer"})
if code != 403 {
t.Fatalf("change owner role should 403, got %d", code)
}
// A 移除 B
code, env = doReq(t, "DELETE", fmt.Sprintf("/api/v1/members/%d", bMember), aAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("remove member failed: %d %s", code, env.Message)
}
// 审计日志存在(登录/建空间/邀请/改角色/移除等)
code, env = doReq(t, "GET", "/api/v1/audit-logs?size=50", aAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("audit list failed: %d %s", code, env.Message)
}
var al struct {
List []models.AuditLog `json:"list"`
Total int64 `json:"total"`
}
_ = json.Unmarshal(env.Data, &al)
if al.Total < 3 {
t.Fatalf("expect >=3 audit logs, got %d", al.Total)
}
}
+247
View File
@@ -0,0 +1,247 @@
package api_test
import (
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/json"
"fmt"
"net/http/httptest"
"strconv"
"strings"
"testing"
"time"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"everypublish/server/internal/models"
"everypublish/shared/proto"
)
func readEnv(t *testing.T, conn *websocket.Conn) *proto.Envelope {
t.Helper()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_, raw, err := conn.Read(ctx)
if err != nil {
t.Fatalf("ws read: %v", err)
}
var env proto.Envelope
if err = json.Unmarshal(raw, &env); err != nil {
t.Fatalf("ws unmarshal: %v", err)
}
return &env
}
func writeEnv(t *testing.T, conn *websocket.Conn, env *proto.Envelope) {
t.Helper()
raw, _ := json.Marshal(env)
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
if err := conn.Write(ctx, websocket.MessageText, raw); err != nil {
t.Fatalf("ws write: %v", err)
}
}
func TestWSAgentFullLoop(t *testing.T) {
resetTestDB()
access, _ := register(t, "agentuser@test.com")
// 1. 配对
code, env := doReq(t, "POST", "/api/v1/agent/pair-code", access, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("pair-code failed: %d %s", code, env.Message)
}
var pc struct {
Code string `json:"code"`
}
_ = json.Unmarshal(env.Data, &pc)
pub, priv, _ := ed25519.GenerateKey(rand.Reader)
pubB64 := base64.StdEncoding.EncodeToString(pub)
code, env = doReq(t, "POST", "/api/v1/agent/pair", "", gin.H{
"code": pc.Code, "deviceName": "test-device", "os": "test", "version": "0.0.1", "publicKey": pubB64,
})
if code != 200 || env.Code != 0 {
t.Fatalf("pair failed: %d %s", code, env.Message)
}
var pd struct {
DeviceID uint64 `json:"deviceId"`
}
_ = json.Unmarshal(env.Data, &pd)
if pd.DeviceID == 0 {
t.Fatal("deviceId missing")
}
// 2. 坏签名 hello 应被拒
srv := httptest.NewServer(testRouter)
defer srv.Close()
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws/agent"
badConn, _, err := websocket.Dial(context.Background(), wsURL, nil)
if err != nil {
t.Fatalf("dial: %v", err)
}
badHello := proto.NewEnvelope("id1", proto.TypeHello, proto.DeviceHello{
DeviceID: strconv.FormatUint(pd.DeviceID, 10),
Nonce: "n1",
TS: time.Now().UnixMilli(),
Sig: base64.StdEncoding.EncodeToString([]byte("bad-sig")),
Version: "0.0.1",
})
rawBad, _ := json.Marshal(badHello)
_ = badConn.Write(context.Background(), websocket.MessageText, rawBad)
ctx2, cancel2 := context.WithTimeout(context.Background(), 3*time.Second)
_, _, err = badConn.Read(ctx2)
cancel2()
if err == nil {
t.Fatal("bad signature should be rejected")
}
// 3. 正常握手
conn, _, err := websocket.Dial(context.Background(), wsURL, nil)
if err != nil {
t.Fatalf("dial: %v", err)
}
nonce := "agent-test-nonce"
ts := time.Now().UnixMilli()
sig := ed25519.Sign(priv, []byte(nonce+"|"+strconv.FormatInt(ts, 10)))
writeEnv(t, conn, proto.NewEnvelope("id2", proto.TypeHello, proto.DeviceHello{
DeviceID: strconv.FormatUint(pd.DeviceID, 10),
Nonce: nonce,
TS: ts,
Sig: base64.StdEncoding.EncodeToString(sig),
Version: "0.0.1",
}))
ackEnv := readEnv(t, conn)
if ackEnv.Type != proto.TypeHelloAck {
t.Fatalf("expect hello.ack, got %s", ackEnv.Type)
}
var ack proto.HelloAck
_ = json.Unmarshal(ackEnv.Payload, &ack)
if ack.SessionID == "" || ack.ServerNonce == "" {
t.Fatal("hello.ack missing fields")
}
// 4. 心跳 → heartbeat.ack
writeEnv(t, conn, proto.NewEnvelope("id3", proto.TypeHeartbeat, proto.Heartbeat{DeviceID: "1", TS: time.Now().UnixMilli()}))
hbAck := readEnv(t, conn)
if hbAck.Type != proto.TypeHeartbeatAck {
t.Fatalf("expect heartbeat.ack, got %s", hbAck.Type)
}
// 5. 账号绑定 → challenge.new → solve → active
code, env = doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "douyin", "accountName": "抖音号01"})
if code != 200 {
t.Fatalf("create account failed: %d", code)
}
var acc models.Account
_ = json.Unmarshal(env.Data, &acc)
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/accounts/%d/bind", acc.ID), access, nil)
if code != 200 {
t.Fatalf("bind failed: %d", code)
}
chEnv := readEnv(t, conn)
if chEnv.Type != proto.TypeChallenge {
t.Fatalf("expect challenge.new, got %s", chEnv.Type)
}
var ch proto.Challenge
_ = json.Unmarshal(chEnv.Payload, &ch)
if ch.AccountID == "" {
t.Fatal("challenge missing accountId")
}
writeEnv(t, conn, proto.NewEnvelope("id5", proto.TypeChallengeSolve, proto.ChallengeSolve{ChallengeID: ch.ChallengeID, Value: "fake-token"}))
time.Sleep(200 * time.Millisecond)
code, env = doReq(t, "GET", "/api/v1/accounts?status=active", access, nil)
if code != 200 {
t.Fatalf("accounts list failed: %d", code)
}
var al struct {
List []models.Account `json:"list"`
}
_ = json.Unmarshal(env.Data, &al)
if len(al.List) != 1 || al.List[0].AgentDeviceID != pd.DeviceID {
t.Fatalf("account should be active bound to device, got %+v", al.List)
}
// 6. 任务闭环(含下发延迟测量)+ 浏览器实时事件
browserConn, _, err := websocket.Dial(context.Background(), "ws"+strings.TrimPrefix(srv.URL, "http")+"/ws/browser?token="+access, nil)
if err != nil {
t.Fatalf("browser dial: %v", err)
}
matID, _ := uploadMaterial(t, access, "loop.mp4", []byte("loop-content"))
code, env = doReq(t, "POST", "/api/v1/tasks", access, gin.H{
"title": "实时闭环任务", "content": "正文",
"accountIds": []uint64{acc.ID}, "materialIds": []uint64{matID},
})
if code != 200 {
t.Fatalf("create task failed: %d", code)
}
var taskRow models.Task
_ = json.Unmarshal(env.Data, &taskRow)
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/submit", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("submit failed: %d", code)
}
t0 := time.Now()
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/approve", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("approve failed: %d", code)
}
pushEnv := readEnv(t, conn)
latency := time.Since(t0)
t.Logf("下发延迟 approve→task.push = %s", latency)
if latency > 2*time.Second {
t.Fatalf("dispatch latency too high: %s", latency)
}
if pushEnv.Type != proto.TypeTaskPush {
t.Fatalf("expect task.push, got %s", pushEnv.Type)
}
var push proto.TaskPush
_ = json.Unmarshal(pushEnv.Payload, &push)
if push.TaskID == "" || len(push.MaterialURLs) == 0 {
t.Fatalf("task.push missing fields: %+v", push)
}
writeEnv(t, conn, proto.NewEnvelope(pushEnv.ID, proto.TypeTaskAck, proto.TaskAck{TaskID: push.TaskID, Accept: true}))
writeEnv(t, conn, proto.NewEnvelope("id6", proto.TypeTaskResult, proto.TaskResult{
TaskID: push.TaskID, Status: "success",
PublishedURL: "https://example.com/v/1",
FinishedAt: time.Now().UnixMilli(),
}))
// 浏览器侧应收到三条 task.status(dispatched/running/success),第三条为 success
bEvents := []string{}
var lastPayload struct {
Status string `json:"status"`
}
for i := 0; i < 3; i++ {
e := readEnv(t, browserConn)
bEvents = append(bEvents, e.Type)
_ = json.Unmarshal(e.Payload, &lastPayload)
}
t.Logf("browser events: %v (last=%s)", bEvents, lastPayload.Status)
if bEvents[0] != "task.status" || bEvents[1] != "task.status" || bEvents[2] != "task.status" {
t.Fatalf("browser should receive 3 task.status events, got %v", bEvents)
}
if lastPayload.Status != "success" {
t.Fatalf("last event should be success, got %s", lastPayload.Status)
}
// 落库校验(轮询至 success)
deadline := time.Now().Add(3 * time.Second)
for {
code, env = doReq(t, "GET", fmt.Sprintf("/api/v1/tasks/%d", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("task get failed: %d", code)
}
_ = json.Unmarshal(env.Data, &taskRow)
if taskRow.Status == "success" || time.Now().After(deadline) {
break
}
time.Sleep(20 * time.Millisecond)
}
if taskRow.Status != "success" {
t.Fatalf("expect success, got %s", taskRow.Status)
}
if !strings.Contains(taskRow.PublishedURLs, "https://example.com/v/1") {
t.Fatalf("publishedUrls not recorded: %s", taskRow.PublishedURLs)
}
}
+18
View File
@@ -0,0 +1,18 @@
package auth
import (
"time"
"github.com/golang-jwt/jwt/v5"
)
func jwtRegisteredExpired() jwt.RegisteredClaims {
return jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(-time.Hour)),
Issuer: "everypublish",
}
}
func signExpired(secret string, claims AccessClaims) (string, error) {
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
}
+66
View File
@@ -0,0 +1,66 @@
package auth
import (
"strings"
"testing"
)
func TestHashVerifyPassword(t *testing.T) {
hash, err := HashPassword("everypublish-2026")
if err != nil {
t.Fatalf("hash: %v", err)
}
if !strings.HasPrefix(hash, "$argon2id$") {
t.Fatalf("bad hash format: %s", hash)
}
ok, err := VerifyPassword("everypublish-2026", hash)
if err != nil || !ok {
t.Fatalf("verify should pass: ok=%v err=%v", ok, err)
}
ok, err = VerifyPassword("wrong-pass", hash)
if err != nil || ok {
t.Fatalf("verify should fail: ok=%v err=%v", ok, err)
}
}
func TestJWTAccessRoundTrip(t *testing.T) {
secret := "test-secret"
token, err := IssueAccess(secret, 42, 7, "user", "owner")
if err != nil {
t.Fatalf("issue: %v", err)
}
claims, err := ParseAccess(secret, token)
if err != nil {
t.Fatalf("parse: %v", err)
}
if claims.UID != 42 || claims.WorkspaceID != 7 || claims.MemberRole != "owner" {
t.Fatalf("claims mismatch: %+v", claims)
}
}
func TestJWTAccessExpired(t *testing.T) {
secret := "test-secret"
claims := AccessClaims{UID: 1, RegisteredClaims: jwtRegisteredExpired()}
token, err := signExpired(secret, claims)
if err != nil {
t.Fatalf("sign: %v", err)
}
if _, err = ParseAccess(secret, token); err == nil {
t.Fatal("expired token should fail")
}
}
func TestJWTRefreshRoundTrip(t *testing.T) {
secret := "test-secret"
token, err := IssueRefresh(secret, 9, "jti-1")
if err != nil {
t.Fatalf("issue: %v", err)
}
claims, err := ParseRefresh(secret, token)
if err != nil {
t.Fatalf("parse: %v", err)
}
if claims.UID != 9 || claims.JTI != "jti-1" {
t.Fatalf("claims mismatch: %+v", claims)
}
}
+135
View File
@@ -0,0 +1,135 @@
package auth
import (
"errors"
"time"
"github.com/golang-jwt/jwt/v5"
)
// 过期策略:Access 15 分钟 / Refresh 7 天(旋转)/ Invite 24 小时
const (
AccessTTL = 15 * time.Minute
RefreshTTL = 7 * 24 * time.Hour
InviteTTL = 24 * time.Hour
)
// AccessClaims 访问令牌声明
type AccessClaims struct {
UID uint64 `json:"uid"`
WorkspaceID uint64 `json:"wsid"`
Role string `json:"role"`
MemberRole string `json:"mrole"`
jwt.RegisteredClaims
}
// RefreshClaims 刷新令牌声明
type RefreshClaims struct {
UID uint64 `json:"uid"`
JTI string `json:"jti"`
jwt.RegisteredClaims
}
// InviteClaims 邀请令牌声明
type InviteClaims struct {
WorkspaceID uint64 `json:"wsid"`
Email string `json:"email"`
Role string `json:"role"`
jwt.RegisteredClaims
}
// IssueAccess 签发访问令牌
func IssueAccess(secret string, uid, wsid uint64, role, memberRole string) (string, error) {
now := time.Now()
claims := AccessClaims{
UID: uid, WorkspaceID: wsid, Role: role, MemberRole: memberRole,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(AccessTTL)),
Issuer: "everypublish",
},
}
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
}
// IssueRefresh 签发刷新令牌
func IssueRefresh(secret string, uid uint64, jti string) (string, error) {
now := time.Now()
claims := RefreshClaims{
UID: uid, JTI: jti,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(RefreshTTL)),
Issuer: "everypublish",
},
}
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
}
// IssueInvite 签发邀请令牌
func IssueInvite(secret string, wsid uint64, email, role string) (string, error) {
now := time.Now()
claims := InviteClaims{
WorkspaceID: wsid, Email: email, Role: role,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(InviteTTL)),
Issuer: "everypublish",
},
}
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
}
// ParseAccess 解析访问令牌
func ParseAccess(secret, token string) (*AccessClaims, error) {
claims := &AccessClaims{}
parsed, err := jwt.ParseWithClaims(token, claims, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, errors.New("unexpected signing method")
}
return []byte(secret), nil
})
if err != nil {
return nil, err
}
if !parsed.Valid {
return nil, errors.New("invalid token")
}
return claims, nil
}
// ParseRefresh 解析刷新令牌
func ParseRefresh(secret, token string) (*RefreshClaims, error) {
claims := &RefreshClaims{}
parsed, err := jwt.ParseWithClaims(token, claims, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, errors.New("unexpected signing method")
}
return []byte(secret), nil
})
if err != nil {
return nil, err
}
if !parsed.Valid {
return nil, errors.New("invalid token")
}
return claims, nil
}
// ParseInvite 解析邀请令牌
func ParseInvite(secret, token string) (*InviteClaims, error) {
claims := &InviteClaims{}
parsed, err := jwt.ParseWithClaims(token, claims, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, errors.New("unexpected signing method")
}
return []byte(secret), nil
})
if err != nil {
return nil, err
}
if !parsed.Valid {
return nil, errors.New("invalid token")
}
return claims, nil
}
+56
View File
@@ -0,0 +1,56 @@
package auth
import (
"crypto/rand"
"crypto/subtle"
"encoding/base64"
"fmt"
"strings"
"golang.org/x/crypto/argon2"
)
// Argon2id 参数(OWASP 推荐基线)
const (
argonTime = 1
argonMemory = 64 * 1024
argonThreads = 4
argonKeyLen = 32
saltLen = 16
)
// HashPassword 生成 $argon2id$v=19$m=65536,t=1,p=4$salt$hash
func HashPassword(password string) (string, error) {
salt := make([]byte, saltLen)
if _, err := rand.Read(salt); err != nil {
return "", err
}
hash := argon2.IDKey([]byte(password), salt, argonTime, argonMemory, argonThreads, argonKeyLen)
b64Salt := base64.RawStdEncoding.EncodeToString(salt)
b64Hash := base64.RawStdEncoding.EncodeToString(hash)
return fmt.Sprintf("$argon2id$v=19$m=%d,t=%d,p=%d$%s$%s", argonMemory, argonTime, argonThreads, b64Salt, b64Hash), nil
}
// VerifyPassword 校验密码(常数时间比较)
func VerifyPassword(password, encoded string) (bool, error) {
parts := strings.Split(encoded, "$")
if len(parts) != 6 {
return false, fmt.Errorf("invalid hash format")
}
var memory uint32
var timeCost uint32
var threads uint8
if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &timeCost, &threads); err != nil {
return false, err
}
salt, err := base64.RawStdEncoding.DecodeString(parts[4])
if err != nil {
return false, err
}
want, err := base64.RawStdEncoding.DecodeString(parts[5])
if err != nil {
return false, err
}
got := argon2.IDKey([]byte(password), salt, timeCost, memory, threads, uint32(len(want)))
return subtle.ConstantTimeCompare(got, want) == 1, nil
}
+31
View File
@@ -0,0 +1,31 @@
package cache
import (
"context"
"time"
"github.com/redis/go-redis/v9"
"everypublish/server/internal/config"
)
// Redis 封装(D4 起 asynq 队列同源共用)
type Redis struct {
client *redis.Client
}
// New 建立 Redis 连接
func New(cfg *config.Config) *Redis {
client := redis.NewClient(&redis.Options{Addr: cfg.RedisAddr, Password: cfg.RedisPass, DB: 0})
return &Redis{client: client}
}
// Ping 探测
func (r *Redis) Ping() error {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
return r.client.Ping(ctx).Err()
}
// Client 暴露底层客户端
func (r *Redis) Client() *redis.Client { return r.client }
+32
View File
@@ -0,0 +1,32 @@
package cache
import (
"context"
"fmt"
"time"
)
// SaveRefresh 登记刷新令牌(jti 白名单,用于轮换与登出吊销)
func (r *Redis) SaveRefresh(uid uint64, jti string, ttl time.Duration) error {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
key := fmt.Sprintf("refresh:%d:%s", uid, jti)
return r.client.Set(ctx, key, "1", ttl).Err()
}
// ExistsRefresh 校验刷新令牌是否有效
func (r *Redis) ExistsRefresh(uid uint64, jti string) (bool, error) {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
key := fmt.Sprintf("refresh:%d:%s", uid, jti)
n, err := r.client.Exists(ctx, key).Result()
return n == 1, err
}
// DeleteRefresh 吊销刷新令牌
func (r *Redis) DeleteRefresh(uid uint64, jti string) error {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
key := fmt.Sprintf("refresh:%d:%s", uid, jti)
return r.client.Del(ctx, key).Err()
}
+61
View File
@@ -0,0 +1,61 @@
package config
import (
"fmt"
"os"
)
// Config 服务器配置(环境变量优先,.env 兜底)
type Config struct {
ServerAddr string
MySQLDSN string
RedisAddr string
RedisPass string
JWTSecret string
BaseURL string
PublicBaseURL string // 浏览器侧访问地址(邀请链接等)
StorageDir string
StaticDir string // 前端生产构建目录(空=不托管静态)
MaxUploadBytes int64 // multipart 上传体积上限
}
// Load 从环境变量加载
func Load() *Config {
jwt := getenv("JWT_SECRET", "dev-secret-change-me")
if jwt == "dev-secret-change-me" || jwt == "please-change-me-in-production" || len(jwt) < 16 {
// 生产安全红线:默认/过短密钥可被伪造 JWT(含 admin 提权)。
// 保留 dev 默认值以便本地开发,但在任何环境都打印醒目告警。
println("!! [security] JWT_SECRET 使用默认值或长度不足(<16)。生产环境必须设置强随机 JWT_SECRET,否则任何人都能伪造令牌(admin 提权)。")
}
return &Config{
ServerAddr: getenv("SERVER_ADDR", ":8090"),
MySQLDSN: getenv("MYSQL_DSN", "root:everypublish@tcp(127.0.0.1:3306)/everypublish?charset=utf8mb4&parseTime=True&loc=Local"),
RedisAddr: getenv("REDIS_ADDR", "127.0.0.1:6379"),
RedisPass: getenv("REDIS_PASSWORD", ""),
JWTSecret: jwt,
BaseURL: getenv("BASE_URL", "http://127.0.0.1:8090"),
PublicBaseURL: getenv("PUBLIC_BASE_URL", getenv("BASE_URL", "http://127.0.0.1:8090")),
StorageDir: getenv("STORAGE_DIR", "./data/materials"),
StaticDir: getenv("STATIC_DIR", "../apps/web/dist"),
MaxUploadBytes: getenvInt64("MAX_UPLOAD_BYTES", 2<<30),
}
}
func getenvInt64(k string, def int64) int64 {
v := os.Getenv(k)
if v == "" {
return def
}
var n int64
if _, err := fmt.Sscanf(v, "%d", &n); err != nil || n <= 0 {
return def
}
return n
}
func getenv(k, def string) string {
if v := os.Getenv(k); v != "" {
return v
}
return def
}
+53
View File
@@ -0,0 +1,53 @@
package db
import (
"fmt"
"time"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
gormmysql "gorm.io/driver/mysql"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// Connect 建立 MySQL 连接并 Ping(重试 10 次,适配容器冷启动)
func Connect(cfg *config.Config) (*gorm.DB, error) {
db, err := gorm.Open(gormmysql.Open(cfg.MySQLDSN), &gorm.Config{Logger: logger.Default.LogMode(logger.Warn)})
if err != nil {
return nil, err
}
sqlDB, err := db.DB()
if err != nil {
return nil, err
}
sqlDB.SetMaxOpenConns(50)
sqlDB.SetMaxIdleConns(10)
sqlDB.SetConnMaxLifetime(time.Hour)
for i := 0; i < 10; i++ {
if err = sqlDB.Ping(); err == nil {
return db, nil
}
time.Sleep(time.Second)
}
return nil, fmt.Errorf("mysql ping failed: %w", err)
}
// AutoMigrate 开发期建表;生产走 golang-migrate(db/migrations)
func AutoMigrate(db *gorm.DB) error {
return db.AutoMigrate(
&models.User{},
&models.Workspace{},
&models.Member{},
&models.Account{},
&models.Material{},
&models.Task{},
&models.AgentDevice{},
&models.Challenge{},
&models.AuditLog{},
&models.Notification{},
&models.PairingCode{},
&models.TransferToken{},
)
}
+166
View File
@@ -0,0 +1,166 @@
package models
import "time"
// Base 通用字段
type Base struct {
ID uint64 `gorm:"primaryKey;autoIncrement" json:"id"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
}
// User 网站用户
type User struct {
Base
Email string `gorm:"uniqueIndex;size:191" json:"email"`
PasswordHash string `gorm:"size:255" json:"-"`
Nickname string `gorm:"size:64" json:"nickname"`
Role string `gorm:"size:16;default:user" json:"role"`
TOTPSecret string `gorm:"size:64" json:"-"`
TOTPEnabled bool `gorm:"default:false" json:"totpEnabled"`
Status string `gorm:"size:16;default:active" json:"status"`
}
// Workspace 工作空间(租户)
type Workspace struct {
Base
Name string `gorm:"size:128" json:"name"`
OwnerID uint64 `gorm:"index" json:"ownerId"`
Plan string `gorm:"size:32;default:free" json:"plan"`
Status string `gorm:"size:16;default:active" json:"status"`
}
// Member 工作空间成员(角色:owner/admin/operator/reviewer/viewer)
type Member struct {
Base
WorkspaceID uint64 `gorm:"uniqueIndex:uk_ws_user" json:"workspaceId"`
UserID uint64 `gorm:"uniqueIndex:uk_ws_user" json:"userId"`
Role string `gorm:"size:16;default:viewer" json:"role"`
}
// Account 平台账号(凭据仅在 Agent 本机保险库,服务器零凭据)
type Account struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Platform string `gorm:"size:32;index" json:"platform"`
AccountName string `gorm:"size:128" json:"accountName"`
Remark string `gorm:"size:255" json:"remark"`
AvatarURL string `gorm:"size:512" json:"avatarUrl"`
Status string `gorm:"size:16;default:unbound" json:"status"`
Health int `gorm:"default:100" json:"health"`
AgentDeviceID uint64 `gorm:"index" json:"agentDeviceId"`
IPProfile string `gorm:"size:191" json:"ipProfile"`
LastActiveAt *time.Time `json:"lastActiveAt"`
LastError string `gorm:"size:512" json:"lastError"`
}
// Material 素材
type Material struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Name string `gorm:"size:255" json:"name"`
Kind string `gorm:"size:16;default:video" json:"kind"`
Size int64 `json:"size"`
SHA256 string `gorm:"size:64;index" json:"sha256"`
Mime string `gorm:"size:128" json:"mime"`
StorageKey string `gorm:"size:512" json:"-"`
Group string `gorm:"size:64" json:"group"`
Tags string `gorm:"type:text" json:"tags"`
Status string `gorm:"size:16;default:ready" json:"status"`
CreatedBy uint64 `json:"createdBy"`
}
// Task 发布任务(状态机:draft→pending_review→approved→queued→dispatched→running→success/failed→retrying;可 cancelled/suspended/rejected)
type Task struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Title string `gorm:"size:255" json:"title"`
Content string `gorm:"type:text" json:"content"`
Tags string `gorm:"type:text" json:"tags"`
AccountIDs string `gorm:"type:text" json:"accountIds"`
MaterialIDs string `gorm:"type:text" json:"materialIds"`
ScheduleAt *time.Time `gorm:"index" json:"scheduleAt"`
Status string `gorm:"size:32;index;default:draft" json:"status"`
Priority int `gorm:"default:5" json:"priority"`
RetryCount int `gorm:"default:0" json:"retryCount"`
MaxRetry int `gorm:"default:3" json:"maxRetry"`
ErrorMessage string `gorm:"size:1024" json:"errorMessage"`
PublishedAt *time.Time `json:"publishedAt"`
PublishedURLs string `gorm:"type:text" json:"publishedUrls"`
Receipts string `gorm:"type:text" json:"receipts"`
CreatedBy uint64 `json:"createdBy"`
ReviewedBy uint64 `json:"reviewedBy"`
ReviewNote string `gorm:"size:512" json:"reviewNote"`
}
// AgentDevice 客户端设备(配对时注册 Ed25519 公钥)
type AgentDevice struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Name string `gorm:"size:128" json:"name"`
OS string `gorm:"size:32" json:"os"`
Version string `gorm:"size:32" json:"version"`
PublicKey string `gorm:"type:text" json:"-"`
Status string `gorm:"size:16;default:offline" json:"status"`
IP string `gorm:"size:64" json:"ip"`
Revoked bool `gorm:"default:false" json:"revoked"`
LastSeenAt *time.Time `json:"lastSeenAt"`
PairedAt time.Time `json:"pairedAt"`
}
// Challenge 二次验证挑战(qr/sms/confirm/captcha/pending;超时挂起,永不暴力破解)
type Challenge struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
AccountID uint64 `gorm:"index" json:"accountId"`
DeviceID uint64 `gorm:"index" json:"deviceId"`
Platform string `gorm:"size:32" json:"platform"`
Kind string `gorm:"size:16" json:"kind"`
Status string `gorm:"size:16;default:active" json:"status"`
QRToken string `gorm:"size:512" json:"qrToken"`
QRURL string `gorm:"size:1024" json:"qrUrl"`
Prompt string `gorm:"size:512" json:"prompt"`
Payload string `gorm:"type:text" json:"payload"`
ExpiresAt time.Time `gorm:"index" json:"expiresAt"`
SolvedAt *time.Time `json:"solvedAt"`
}
// AuditLog 审计日志
type AuditLog struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
UserID uint64 `gorm:"index" json:"userId"`
Action string `gorm:"size:64;index" json:"action"`
Resource string `gorm:"size:128" json:"resource"`
Detail string `gorm:"type:text" json:"detail"`
IP string `gorm:"size:64" json:"ip"`
}
// Notification 站内通知
type Notification struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
UserID uint64 `gorm:"index" json:"userId"`
Kind string `gorm:"size:32" json:"kind"`
Title string `gorm:"size:255" json:"title"`
Content string `gorm:"type:text" json:"content"`
Read bool `gorm:"default:false" json:"read"`
}
// PairingCode 配对码(一次性,5 分钟过期)
type PairingCode struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Code string `gorm:"size:12;uniqueIndex" json:"code"`
ExpiresAt time.Time `json:"expiresAt"`
Used bool `gorm:"default:false" json:"used"`
DeviceID uint64 `json:"deviceId"`
}
// TransferToken 素材直链下载 token(短时签名)
type TransferToken struct {
Base
MaterialID uint64 `gorm:"index" json:"materialId"`
Token string `gorm:"size:191;uniqueIndex" json:"token"`
ExpiresAt time.Time `json:"expiresAt"`
}
+87
View File
@@ -0,0 +1,87 @@
package task
// 任务状态
const (
Draft = "draft"
PendingReview = "pending_review"
Rejected = "rejected"
Queued = "queued"
Dispatched = "dispatched"
Running = "running"
Success = "success"
Failed = "failed"
Suspended = "suspended"
Cancelled = "cancelled"
)
// 动作
type Action string
const (
ActionSubmit Action = "submit"
ActionApprove Action = "approve"
ActionReject Action = "reject"
ActionResubmit Action = "resubmit"
ActionDispatch Action = "dispatch"
ActionStart Action = "start"
ActionSuccess Action = "success"
ActionFail Action = "fail"
ActionRetry Action = "retry"
ActionSuspend Action = "suspend"
ActionResume Action = "resume"
ActionCancel Action = "cancel"
)
// transitions 状态转移表
var transitions = map[string]map[Action]string{
Draft: {
ActionSubmit: PendingReview,
ActionCancel: Cancelled,
},
PendingReview: {
ActionApprove: Queued,
ActionReject: Rejected,
ActionCancel: Cancelled,
},
Rejected: {
ActionResubmit: PendingReview,
ActionCancel: Cancelled,
},
Queued: {
ActionDispatch: Dispatched,
ActionCancel: Cancelled,
},
Dispatched: {
ActionStart: Running,
ActionFail: Failed,
ActionCancel: Cancelled,
},
Running: {
ActionSuccess: Success,
ActionFail: Failed,
ActionSuspend: Suspended,
ActionCancel: Cancelled,
},
Failed: {
ActionRetry: Queued,
ActionCancel: Cancelled,
},
Suspended: {
ActionResume: Queued,
ActionCancel: Cancelled,
},
Success: {},
Cancelled: {},
}
// Next 校验并返回下一状态;非法转移返回 ok=false
func Next(current string, a Action) (string, bool) {
next, ok := transitions[current][a]
return next, ok
}
// Can 当前状态是否接受该动作
func Can(current string, a Action) bool {
_, ok := transitions[current][a]
return ok
}
+82
View File
@@ -0,0 +1,82 @@
package task
import "testing"
func TestHappyPath(t *testing.T) {
path := []struct {
action Action
want string
}{
{ActionSubmit, PendingReview},
{ActionApprove, Queued},
{ActionDispatch, Dispatched},
{ActionStart, Running},
{ActionSuccess, Success},
}
cur := Draft
for _, step := range path {
next, ok := Next(cur, step.action)
if !ok || next != step.want {
t.Fatalf("step %s from %s: got %s ok=%v, want %s", step.action, cur, next, ok, step.want)
}
cur = next
}
}
func TestFailureRetry(t *testing.T) {
cur, ok := Next(Draft, ActionSubmit)
if !ok {
t.Fatal("submit failed")
}
cur, ok = Next(cur, ActionApprove)
if !ok {
t.Fatal("approve failed")
}
cur, ok = Next(cur, ActionDispatch)
if !ok {
t.Fatal("dispatch failed")
}
cur, ok = Next(cur, ActionStart)
if !ok {
t.Fatal("start failed")
}
cur, ok = Next(cur, ActionFail)
if !ok || cur != Failed {
t.Fatalf("fail: got %s ok=%v", cur, ok)
}
cur, ok = Next(cur, ActionRetry)
if !ok || cur != Queued {
t.Fatalf("retry: got %s ok=%v", cur, ok)
}
}
func TestSuspendResume(t *testing.T) {
cur, _ := Next(Running, ActionSuspend)
if cur != Suspended {
t.Fatalf("suspend: got %s", cur)
}
cur, ok := Next(cur, ActionResume)
if !ok || cur != Queued {
t.Fatalf("resume: got %s ok=%v", cur, ok)
}
}
func TestIllegalTransitions(t *testing.T) {
cases := []struct {
cur string
action Action
}{
{Draft, ActionApprove},
{PendingReview, ActionStart},
{Queued, ActionSuccess},
{Success, ActionRetry},
{Cancelled, ActionResubmit},
{Running, ActionApprove},
{Dispatched, ActionSuspend},
}
for _, c := range cases {
if _, ok := Next(c.cur, c.action); ok {
t.Fatalf("illegal transition should fail: %s -> %s", c.cur, c.action)
}
}
}
+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
}