server: 平台抽象层+bilibili、服务端加密凭据库、角色中间件;web: 精简页面/路由、macOS 客户端(Swift)与多份方案文档;移除误入库的编译产物

This commit is contained in:
Qiufeng
2026-08-21 10:53:32 +08:00
parent 0daa9782c9
commit 3585c39bab
103 changed files with 5842 additions and 3197 deletions
+47 -10
View File
@@ -46,7 +46,7 @@ func uploadMaterial(t *testing.T, access string, filename string, content []byte
}
var data struct {
Material models.Material `json:"material"`
Dedup bool `json:"dedup"`
Dedup bool `json:"dedup"`
}
_ = json.Unmarshal(env.Data, &data)
if data.Dedup {
@@ -88,8 +88,8 @@ func TestAccountChallengeFlow(t *testing.T) {
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))
if len(cl.List) != 1 || cl.List[0].Kind != "qr" || cl.List[0].QRURL == "" {
t.Fatalf("expect 1 web QR challenge, got %d", len(cl.List))
}
// 挂起 → 重发
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/challenges/%d/suspend", cl.List[0].ID), access, nil)
@@ -121,7 +121,7 @@ func TestMaterialAndTaskFlow(t *testing.T) {
_ = json.Unmarshal(w.Body.Bytes(), &env)
var data struct {
Material models.Material `json:"material"`
Dedup bool `json:"dedup"`
Dedup bool `json:"dedup"`
}
_ = json.Unmarshal(env.Data, &data)
if !data.Dedup || data.Material.ID != matID {
@@ -148,10 +148,45 @@ func TestMaterialAndTaskFlow(t *testing.T) {
if code != 404 {
t.Fatalf("one-time token should 404 on reuse, got %d", code)
}
accountRespCode, accountEnv := doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "douyin", "accountName": "任务测试账号"})
if accountRespCode != 200 || accountEnv.Code != 0 {
t.Fatalf("create task account failed: %d %s", accountRespCode, accountEnv.Message)
}
var taskAccount models.Account
_ = json.Unmarshal(accountEnv.Data, &taskAccount)
bindCode, bindEnv := doReq(t, "POST", fmt.Sprintf("/api/v1/accounts/%d/bind", taskAccount.ID), access, nil)
if bindCode != 200 || bindEnv.Code != 0 {
t.Fatalf("task account bind failed: %d %s", bindCode, bindEnv.Message)
}
var challengeList struct {
List []models.Challenge `json:"list"`
}
_, challengeEnv := doReq(t, "GET", "/api/v1/challenges", access, nil)
_ = json.Unmarshal(challengeEnv.Data, &challengeList)
if len(challengeList.List) == 0 {
t.Fatal("task account challenge missing")
}
challengeCode, challengeEnv := doReq(t, "POST", fmt.Sprintf("/api/v1/challenges/%d/solve", challengeList.List[0].ID), access, gin.H{"value": "web-test-login"})
if challengeCode != 200 || challengeEnv.Code != 0 {
t.Fatalf("task account challenge solve failed: %d %s", challengeCode, challengeEnv.Message)
}
// 后端不能接受尚未登录的账号作为发布目标,即使客户端绕过了筛选。
unboundCode, unboundEnv := doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "bilibili", "remark": "未登录账号"})
if unboundCode != 200 || unboundEnv.Code != 0 {
t.Fatalf("create unbound account failed: %d %s", unboundCode, unboundEnv.Message)
}
var unbound models.Account
_ = json.Unmarshal(unboundEnv.Data, &unbound)
badTaskCode, _ := doReq(t, "POST", "/api/v1/tasks", access, gin.H{
"title": "不应创建", "content": "", "accountIds": []uint64{unbound.ID}, "materialIds": []uint64{matID},
})
if badTaskCode != 400 {
t.Fatalf("unbound account task should be rejected, got %d", badTaskCode)
}
// 建任务 → 提交 → 驳回 → 重提 → 通过 → 取消
code, env = doReq(t, "POST", "/api/v1/tasks", access, gin.H{
"title": "新品发布", "content": "正文", "tags": []string{"新品"},
"accountIds": []uint64{1}, "materialIds": []uint64{matID},
"accountIds": []uint64{taskAccount.ID}, "materialIds": []uint64{matID},
})
if code != 200 || env.Code != 0 {
t.Fatalf("create task failed: %d %s", code, env.Message)
@@ -189,12 +224,14 @@ func TestMaterialAndTaskFlow(t *testing.T) {
}
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)
if taskRow.Status != "queued" && taskRow.Status != "dispatched" && taskRow.Status != "running" && taskRow.Status != "success" {
t.Fatalf("expect web executor to accept task, 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)
if taskRow.Status == "queued" || taskRow.Status == "dispatched" || taskRow.Status == "running" {
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)
+334 -47
View File
@@ -1,24 +1,36 @@
package handlers
import (
"context"
"errors"
"net/http"
"strconv"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/credentials"
"everypublish/server/internal/models"
"everypublish/server/internal/platform"
"everypublish/server/internal/ws"
"everypublish/shared/proto"
)
// AccountHandler 账号台账处理器(凭据仅在 Agent 本机保险库,服务器零凭据)
// AccountHandler 账号台账处理器(Web-only 本机执行器)。
type AccountHandler struct {
DB *gorm.DB
Hub *ws.Hub
DB *gorm.DB
Hub *ws.Hub
Registry *platform.Registry
Credentials *credentials.Store
ExecutorMode string
PollInterval time.Duration
pollMu sync.Mutex
pollCancels map[uint64]context.CancelFunc
pollGeneration map[uint64]uint64
pollSeq uint64
}
var platforms = map[string]bool{
@@ -27,10 +39,15 @@ var platforms = map[string]bool{
}
type accountReq struct {
Platform string `json:"platform" binding:"required"`
Remark string `json:"remark"`
AgentDeviceID uint64 `json:"agentDeviceId"`
AccountName string `json:"accountName"`
Platform string `json:"platform" binding:"required"`
Remark string `json:"remark"`
AccountName string `json:"accountName"`
AvatarURL string `json:"avatarUrl"`
IPProfile string `json:"ipProfile"`
}
type credentialsReq struct {
Cookies string `json:"cookies" binding:"required"`
}
// List 账号台账(筛选 platform/status)
@@ -57,7 +74,7 @@ func (h *AccountHandler) List(c *gin.Context) {
response.OK(c, gin.H{"list": list, "total": total, "page": page, "size": size})
}
// Create 新增账号(初始 unbound;凭据只登记元数据)
// Create 新增账号(初始 unbound;V1 凭据由本机 Web executor 管理)
func (h *AccountHandler) Create(c *gin.Context) {
var req accountReq
if err := c.ShouldBindJSON(&req); err != nil {
@@ -69,12 +86,14 @@ func (h *AccountHandler) Create(c *gin.Context) {
return
}
account := models.Account{
WorkspaceID: c.GetUint64("wsid"),
Platform: req.Platform,
Remark: req.Remark,
AgentDeviceID: req.AgentDeviceID,
Status: "unbound",
Health: 100,
WorkspaceID: c.GetUint64("wsid"),
Platform: req.Platform,
AccountName: req.AccountName,
Remark: req.Remark,
AvatarURL: req.AvatarURL,
IPProfile: req.IPProfile,
Status: "unbound",
Health: 100,
}
if err := h.DB.Create(&account).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "创建失败")
@@ -91,19 +110,26 @@ func (h *AccountHandler) Update(c *gin.Context) {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req accountReq
var req struct {
Remark string `json:"remark"`
AccountName string `json:"accountName"`
AvatarURL string `json:"avatarUrl"`
IPProfile string `json:"ipProfile"`
}
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
updates := map[string]interface{}{"remark": req.Remark}
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
}
updates := map[string]interface{}{"remark": req.Remark, "avatar_url": req.AvatarURL, "ip_profile": req.IPProfile}
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 {
if err = h.DB.Model(&account).Updates(updates).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
@@ -131,11 +157,14 @@ func (h *AccountHandler) Delete(c *gin.Context) {
response.Fail(c, http.StatusInternalServerError, 5000, "删除失败")
return
}
if h.Credentials != nil {
_ = h.Credentials.Delete(account.WorkspaceID, account.ID)
}
response.Audit(c, "account.delete", "account:"+itoa(id), nil)
response.OK(c, nil)
}
// Bind 发起绑定:创建挑战记录(pending),等待客户端扫码(D4 WSS 联动)
// Bind 发起绑定:创建网页端挑战,由当前服务器上的 Web adapter 继续处理。
func (h *AccountHandler) Bind(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
@@ -147,19 +176,50 @@ func (h *AccountHandler) Bind(c *gin.Context) {
response.Fail(c, http.StatusNotFound, 1004, "账号不存在")
return
}
if account.Status == "active" {
response.Fail(c, http.StatusConflict, 3002, "账号已绑定")
return
if account.Status == "binding" {
var existing models.Challenge
if findErr := h.DB.Where("account_id = ? and workspace_id = ? and status = ?", account.ID, account.WorkspaceID, "active").Order("id desc").First(&existing).Error; findErr == nil {
response.OK(c, gin.H{"challengeId": existing.ID, "account": account.ID, "status": "binding", "qrUrl": existing.QRURL, "qrToken": existing.QRToken, "prompt": existing.Prompt, "expiresAt": existing.ExpiresAt})
return
}
}
var staleChallenges []models.Challenge
_ = h.DB.Where("account_id = ? and workspace_id = ? and status = ?", account.ID, account.WorkspaceID, "active").Find(&staleChallenges).Error
for _, stale := range staleChallenges {
h.cancelLoginPoll(stale.ID)
}
if len(staleChallenges) > 0 {
_ = h.DB.Model(&models.Challenge{}).Where("account_id = ? and workspace_id = ? and status = ?", account.ID, account.WorkspaceID, "active").Update("status", "suspended").Error
}
// A registered adapter owns the login session. The browser receives one QR
// URL and the server polls it; no second client-side pairing code is needed.
var adapter platform.Adapter
if h.ExecutorMode != "mock" && h.Registry != nil {
adapter = h.Registry.Get(account.Platform)
}
var session platform.LoginSession
if adapter != nil {
var beginErr error
session, beginErr = adapter.BeginLogin(c.Request.Context())
if beginErr != nil {
response.Fail(c, http.StatusBadGateway, 2007, adapterErrorMessage(beginErr))
return
}
}
challenge := models.Challenge{
WorkspaceID: c.GetUint64("wsid"),
AccountID: account.ID,
Platform: account.Platform,
Kind: "pending",
Kind: "qr",
Status: "active",
Prompt: "等待客户端发起扫码,稍后此处展示二维码",
Prompt: "请在网页中完成 " + account.Platform + " 登录",
QRToken: session.Token,
QRURL: session.URL,
ExpiresAt: time.Now().Add(30 * time.Minute),
}
if adapter == nil {
challenge.QRURL = "mock://everypublish/login/" + strconv.FormatUint(account.ID, 10)
}
err = h.DB.Transaction(func(tx *gorm.DB) error {
if err = tx.Create(&challenge).Error; err != nil {
return err
@@ -170,26 +230,253 @@ func (h *AccountHandler) Bind(c *gin.Context) {
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)
}
h.Hub.BroadcastWS(c.GetUint64("wsid"), "challenge.status", gin.H{
"challengeId": challenge.ID, "accountId": account.ID, "status": challenge.Status,
"kind": challenge.Kind, "platform": challenge.Platform, "qrUrl": challenge.QRURL,
"qrToken": challenge.QRToken, "prompt": challenge.Prompt, "expiresAt": challenge.ExpiresAt,
})
}
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"})
if adapter != nil {
h.StartLoginPoll(challenge, account, c.GetUint64("uid"), adapter)
}
response.OK(c, gin.H{"challengeId": challenge.ID, "account": account.ID, "status": "binding", "qrUrl": challenge.QRURL, "qrToken": challenge.QRToken, "prompt": challenge.Prompt, "expiresAt": challenge.ExpiresAt})
}
// ImportCredentials accepts a user-supplied platform cookie string. The value
// is written directly to the encrypted adapter store and never returned or
// included in audit details.
func (h *AccountHandler) ImportCredentials(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req credentialsReq
if err = c.ShouldBindJSON(&req); err != nil || len(strings.TrimSpace(req.Cookies)) < 8 || len(req.Cookies) > 64*1024 {
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 h.ExecutorMode == "mock" || h.Registry == nil {
response.Fail(c, http.StatusConflict, 3006, "当前 mock 模式不支持真实凭据导入")
return
}
adapter := h.Registry.Get(account.Platform)
if adapter == nil {
response.Fail(c, http.StatusConflict, 3006, "该平台尚未接入真实适配器")
return
}
if err = adapter.SaveCredentials(account.WorkspaceID, account.ID, strings.TrimSpace(req.Cookies)); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, adapterErrorMessage(err))
return
}
now := time.Now()
if err = h.DB.Model(&account).Updates(map[string]interface{}{"status": "active", "last_active_at": &now, "last_error": ""}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "账号状态更新失败")
return
}
var activeChallenges []models.Challenge
_ = h.DB.Where("account_id = ? and workspace_id = ? and status = ?", account.ID, account.WorkspaceID, "active").Find(&activeChallenges).Error
for _, challenge := range activeChallenges {
h.cancelLoginPoll(challenge.ID)
}
_ = h.DB.Model(&models.Challenge{}).Where("account_id = ? and workspace_id = ? and status = ?", account.ID, account.WorkspaceID, "active").Update("status", "suspended").Error
_ = Notify(h.DB, account.WorkspaceID, []uint64{c.GetUint64("uid")}, "challenge", "账号凭据已导入", "账号已恢复可用")
if h.Hub != nil {
h.Hub.BroadcastWS(account.WorkspaceID, "account.status", gin.H{"accountId": account.ID, "status": "active", "platform": account.Platform, "prompt": "凭据已导入"})
}
response.Audit(c, "account.credentials_import", "account:"+itoa(account.ID), gin.H{"platform": account.Platform})
response.OK(c, gin.H{"accountId": account.ID, "status": "active", "lastActiveAt": now})
}
func (h *AccountHandler) Check(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 h.Registry == nil || h.ExecutorMode == "mock" {
response.Fail(c, http.StatusConflict, 3006, "当前模式不支持真实账号检查")
return
}
checker, ok := h.Registry.Get(account.Platform).(platform.CredentialChecker)
if !ok {
response.Fail(c, http.StatusConflict, 3006, "该平台尚未提供账号检查")
return
}
if err = checker.CheckCredentials(c.Request.Context(), account.WorkspaceID, account.ID); err != nil {
message := adapterErrorMessage(err)
_ = h.DB.Model(&account).Updates(map[string]interface{}{"status": "expired", "last_error": message}).Error
response.Fail(c, http.StatusBadGateway, 2007, message)
return
}
now := time.Now()
_ = h.DB.Model(&account).Updates(map[string]interface{}{"status": "active", "last_active_at": &now, "last_error": ""}).Error
response.Audit(c, "account.credentials_check", "account:"+itoa(account.ID), gin.H{"platform": account.Platform})
response.OK(c, gin.H{"accountId": account.ID, "status": "active", "lastActiveAt": now})
}
// Unbind 解绑账号,清除本机执行器关联但保留账号台账记录。
func (h *AccountHandler) Unbind(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 err = h.DB.Model(&account).Updates(map[string]interface{}{
"status": "unbound", "account_name": "", "last_error": "",
}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "解绑失败")
return
}
_ = h.DB.Model(&models.Challenge{}).Where("account_id = ? and workspace_id = ? and status = ?", account.ID, account.WorkspaceID, "active").Update("status", "suspended").Error
if h.Credentials != nil {
_ = h.Credentials.Delete(account.WorkspaceID, account.ID)
}
response.Audit(c, "account.unbind", "account:"+itoa(id), nil)
response.OK(c, nil)
}
func (h *AccountHandler) pollLogin(ctx context.Context, challengeID uint64, account models.Account, userID uint64, adapter platform.Adapter) {
interval := h.PollInterval
if interval <= 0 {
interval = 2 * time.Second
}
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
var challenge models.Challenge
if err := h.DB.Where("id = ? and workspace_id = ?", challengeID, account.WorkspaceID).First(&challenge).Error; err != nil || challenge.Status != "active" {
return
}
if time.Now().After(challenge.ExpiresAt) {
h.finishLogin(challenge, account, userID, "expired", "登录二维码已过期")
return
}
poll, err := adapter.PollLogin(ctx, challenge.QRToken)
if err == nil {
switch poll.State {
case platform.LoginPending, platform.LoginScanned:
prompt := poll.Prompt
if prompt == "" {
prompt = challenge.Prompt
}
_ = h.DB.Model(&challenge).Update("prompt", prompt).Error
h.broadcastChallenge(challenge, string(poll.State), prompt)
case platform.LoginConfirmed:
if saveErr := adapter.SaveCredentials(account.WorkspaceID, account.ID, poll.Cookies); saveErr != nil {
h.finishLogin(challenge, account, userID, "suspended", adapterErrorMessage(saveErr))
return
}
now := time.Now()
if h.DB.Model(&challenge).Updates(map[string]interface{}{"status": "solved", "solved_at": &now, "prompt": "登录成功"}).Error != nil {
return
}
_ = h.DB.Model(&models.Account{}).Where("id = ? and workspace_id = ?", account.ID, account.WorkspaceID).Updates(map[string]interface{}{"status": "active", "last_active_at": &now, "last_error": ""}).Error
_ = Notify(h.DB, account.WorkspaceID, []uint64{userID}, "challenge", "登录挑战已完成", "账号已恢复可用")
h.broadcastChallenge(challenge, "solved", "登录成功")
return
}
} else if errors.Is(err, context.Canceled) {
return
} else if errors.Is(err, context.DeadlineExceeded) {
h.finishLogin(challenge, account, userID, "expired", "登录等待已超时")
return
} else if platform.CategoryOf(err) != platform.Network {
h.finishLogin(challenge, account, userID, "suspended", adapterErrorMessage(err))
return
}
select {
case <-ctx.Done():
if errors.Is(ctx.Err(), context.DeadlineExceeded) {
h.finishLogin(challenge, account, userID, "expired", "登录等待已超时")
}
return
case <-ticker.C:
}
}
}
// StartLoginPoll is shared with challenge resend so every real QR session has
// exactly one backend worker responsible for polling and credential persistence.
func (h *AccountHandler) StartLoginPoll(challenge models.Challenge, account models.Account, userID uint64, adapter platform.Adapter) {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Minute)
h.pollMu.Lock()
if h.pollCancels == nil {
h.pollCancels = make(map[uint64]context.CancelFunc)
}
if h.pollGeneration == nil {
h.pollGeneration = make(map[uint64]uint64)
}
if old := h.pollCancels[challenge.ID]; old != nil {
old()
}
h.pollSeq++
generation := h.pollSeq
h.pollCancels[challenge.ID] = cancel
h.pollGeneration[challenge.ID] = generation
h.pollMu.Unlock()
go func() {
h.pollLogin(ctx, challenge.ID, account, userID, adapter)
h.pollMu.Lock()
if h.pollGeneration[challenge.ID] == generation {
delete(h.pollCancels, challenge.ID)
delete(h.pollGeneration, challenge.ID)
}
h.pollMu.Unlock()
}()
}
func (h *AccountHandler) cancelLoginPoll(challengeID uint64) {
h.pollMu.Lock()
cancel := h.pollCancels[challengeID]
h.pollMu.Unlock()
if cancel != nil {
cancel()
}
}
func (h *AccountHandler) finishLogin(challenge models.Challenge, account models.Account, userID uint64, status, message string) {
_ = h.DB.Model(&challenge).Updates(map[string]interface{}{"status": status, "prompt": message}).Error
accountStatus := "expired"
if status == "suspended" {
accountStatus = "suspended"
}
_ = h.DB.Model(&models.Account{}).Where("id = ? and workspace_id = ?", account.ID, account.WorkspaceID).Updates(map[string]interface{}{"status": accountStatus, "last_error": message}).Error
_ = Notify(h.DB, account.WorkspaceID, []uint64{userID}, "challenge", "登录需要处理", message)
h.broadcastChallenge(challenge, status, message)
}
func (h *AccountHandler) broadcastChallenge(challenge models.Challenge, status, prompt string) {
if h.Hub == nil {
return
}
h.Hub.BroadcastWS(challenge.WorkspaceID, "challenge.status", gin.H{
"challengeId": challenge.ID, "accountId": challenge.AccountID, "status": status,
"kind": challenge.Kind, "platform": challenge.Platform, "prompt": prompt,
})
}
func adapterErrorMessage(err error) string {
var pe *platform.Error
if errors.As(err, &pe) && strings.TrimSpace(pe.Message) != "" {
return pe.Message
}
return "平台登录暂时不可用,请稍后重试"
}
@@ -0,0 +1,112 @@
package handlers
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"sync"
"testing"
"time"
"everypublish/server/internal/models"
"everypublish/server/internal/platform"
"github.com/gin-gonic/gin"
gormsqlite "gorm.io/driver/sqlite"
"gorm.io/gorm"
)
type loginTestAdapter struct {
mu sync.Mutex
polls int
saved string
}
func (a *loginTestAdapter) Platform() string { return "test-platform" }
func (a *loginTestAdapter) BeginLogin(context.Context) (platform.LoginSession, error) {
return platform.LoginSession{Token: "token-1", URL: "https://qr.test/1"}, nil
}
func (a *loginTestAdapter) PollLogin(context.Context, string) (platform.LoginPoll, error) {
a.mu.Lock()
defer a.mu.Unlock()
a.polls++
if a.polls == 1 {
return platform.LoginPoll{State: platform.LoginPending, Prompt: "等待扫码"}, nil
}
return platform.LoginPoll{State: platform.LoginConfirmed, Cookies: "cookie-value"}, nil
}
func (a *loginTestAdapter) SaveCredentials(_ uint64, _ uint64, cookies string) error {
a.mu.Lock()
a.saved = cookies
a.mu.Unlock()
return nil
}
func (a *loginTestAdapter) Publish(context.Context, platform.PublishInput) (platform.PublishResult, error) {
return platform.PublishResult{}, nil
}
func TestAccountBindPollsAdapterAndActivatesAccount(t *testing.T) {
db, err := gorm.Open(gormsqlite.Open(t.TempDir()+"/account.db"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err = db.AutoMigrate(&models.Account{}, &models.Challenge{}, &models.Notification{}, &models.AuditLog{}); err != nil {
t.Fatal(err)
}
account := models.Account{WorkspaceID: 9, Platform: "test-platform", Status: "unbound"}
if err = db.Create(&account).Error; err != nil {
t.Fatal(err)
}
adapter := &loginTestAdapter{}
registry := platform.NewRegistry()
if err = registry.Register(adapter); err != nil {
t.Fatal(err)
}
h := &AccountHandler{DB: db, Registry: registry, ExecutorMode: "web", PollInterval: 5 * time.Millisecond}
gin.SetMode(gin.TestMode)
req := httptest.NewRequest(http.MethodPost, "/accounts/1/bind", nil)
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Request = req
ctx.Params = gin.Params{{Key: "id", Value: "1"}}
ctx.Set("wsid", uint64(9))
ctx.Set("uid", uint64(3))
h.Bind(ctx)
if recorder.Code != http.StatusOK {
t.Fatalf("bind status=%d body=%s", recorder.Code, recorder.Body.String())
}
var envelope struct {
Data struct {
ChallengeID uint64 `json:"challengeId"`
QRURL string `json:"qrUrl"`
} `json:"data"`
}
if err = json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil {
t.Fatal(err)
}
if envelope.Data.ChallengeID == 0 || envelope.Data.QRURL != "https://qr.test/1" {
t.Fatalf("unexpected bind response: %s", recorder.Body.String())
}
deadline := time.Now().Add(time.Second)
for time.Now().Before(deadline) {
var got models.Account
_ = db.First(&got, account.ID).Error
if got.Status == "active" {
adapter.mu.Lock()
saved := adapter.saved
adapter.mu.Unlock()
if saved != "cookie-value" {
t.Fatalf("credentials not saved: %q", saved)
}
var challenge models.Challenge
_ = db.First(&challenge, envelope.Data.ChallengeID).Error
if challenge.Status != "solved" {
t.Fatalf("challenge status=%s", challenge.Status)
}
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatal("adapter login worker did not activate account")
}
+3 -5
View File
@@ -20,12 +20,11 @@ type AdminHandler struct {
type tenantRow struct {
models.Workspace
OwnerEmail string `json:"ownerEmail"`
MemberCount int64 `json:"memberCount"`
DeviceCount int64 `json:"deviceCount"`
OwnerEmail string `json:"ownerEmail"`
MemberCount int64 `json:"memberCount"`
}
// Tenants 租户列表(含 owner 邮箱与成员/设备数)
// Tenants 租户列表(含 owner 邮箱与成员数)
func (h *AdminHandler) Tenants(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
@@ -51,7 +50,6 @@ func (h *AdminHandler) Tenants(c *gin.Context) {
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})
+47 -5
View File
@@ -1,10 +1,13 @@
package handlers
import (
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"errors"
"net/http"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
@@ -52,18 +55,45 @@ func (h *AgentHandler) PairCode(c *gin.Context) {
// Pair 设备配对(无鉴权,凭一次性码;注册 Ed25519 公钥)
func (h *AgentHandler) Pair(c *gin.Context) {
var req struct {
Code string `json:"code" binding:"required"`
DeviceName string `json:"deviceName" binding:"required"`
Code string `json:"code"`
DeviceName string `json:"deviceName"`
OS string `json:"os"`
Version string `json:"version"`
PublicKey string `json:"publicKey" binding:"required"`
PublicKey string `json:"publicKey"`
}
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
response.Fail(c, http.StatusBadRequest, 1001, "请求体必须是 JSON")
return
}
req.Code = strings.ToUpper(strings.TrimSpace(req.Code))
req.DeviceName = strings.TrimSpace(req.DeviceName)
req.PublicKey = strings.TrimSpace(req.PublicKey)
if req.Code == "" {
response.Fail(c, http.StatusBadRequest, 1001, "连接码不能为空")
return
}
if !validPairCode(req.Code) {
response.Fail(c, http.StatusBadRequest, 1001, "配对码应为 6 位大写字母或数字")
return
}
if req.PublicKey == "" {
response.Fail(c, http.StatusBadRequest, 1001, "设备公钥不能为空,请重启客户端后重试")
return
}
publicKey, keyErr := base64.StdEncoding.DecodeString(req.PublicKey)
if keyErr != nil || len(publicKey) != ed25519.PublicKeySize {
response.Fail(c, http.StatusBadRequest, 1001, "设备公钥格式不合法")
return
}
if req.DeviceName == "" {
req.DeviceName = "macOS Agent"
}
var pc models.PairingCode
if err := h.DB.Where("code = ? and used = ?", req.Code, false).First(&pc).Error; err != nil {
if !errors.Is(err, gorm.ErrRecordNotFound) {
response.Fail(c, http.StatusInternalServerError, 5000, "查询配对码失败")
return
}
response.Fail(c, http.StatusBadRequest, 2006, "配对码无效")
return
}
@@ -84,7 +114,7 @@ func (h *AgentHandler) Pair(c *gin.Context) {
// 原子占用:仅当仍 unused 且未过期时置 used=true;RowsAffected==1 才视为抢到,
// 防止并发用同一码配出多台设备。
res := tx.Model(&models.PairingCode{}).
Where("id = ? and used = ?", pc.ID, false).
Where("id = ? and used = ? and expires_at > ?", pc.ID, false, time.Now()).
Update("used", true)
if res.Error != nil {
return res.Error
@@ -112,6 +142,18 @@ func (h *AgentHandler) Pair(c *gin.Context) {
response.OK(c, gin.H{"deviceId": device.ID, "workspaceId": device.WorkspaceID})
}
func validPairCode(code string) bool {
if len(code) != 6 {
return false
}
for _, r := range code {
if !strings.ContainsRune(pairAlphabet, r) {
return false
}
}
return true
}
// Devices 设备列表(在线状态来自 Hub,此处以 DB 状态兜底)
func (h *AgentHandler) Devices(c *gin.Context) {
var list []models.AgentDevice
@@ -0,0 +1,24 @@
package handlers
import "testing"
func TestValidPairCode(t *testing.T) {
tests := []struct {
name string
code string
want bool
}{
{name: "valid", code: "AB23CD", want: true},
{name: "ambiguous characters rejected", code: "AB01CD", want: false},
{name: "lowercase rejected before normalization", code: "ab12cd", want: false},
{name: "wrong length", code: "AB12C", want: false},
{name: "unicode rejected", code: "AB12CD", want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := validPairCode(tt.code); got != tt.want {
t.Fatalf("validPairCode(%q) = %v, want %v", tt.code, got, tt.want)
}
})
}
}
+74 -2
View File
@@ -2,6 +2,7 @@ package handlers
import (
"net/http"
"strconv"
"strings"
"github.com/gin-gonic/gin"
@@ -42,6 +43,11 @@ type logoutReq struct {
RefreshToken string `json:"refreshToken"`
}
type changePasswordReq struct {
CurrentPassword string `json:"currentPassword" binding:"required"`
NewPassword string `json:"newPassword" binding:"required,min=8"`
}
// Register 注册:创建用户 + 默认工作空间 + owner 成员
func (h *AuthHandler) Register(c *gin.Context) {
var req registerReq
@@ -177,8 +183,13 @@ func (h *AuthHandler) Refresh(c *gin.Context) {
response.Fail(c, http.StatusUnauthorized, 2003, "用户不存在")
return
}
workspaceID := claims.WorkspaceID
var member models.Member
if err = h.DB.Where("user_id = ?", user.ID).Order("id asc").First(&member).Error; err != nil {
memberQuery := h.DB.Where("user_id = ?", user.ID)
if workspaceID != 0 {
memberQuery = memberQuery.Where("workspace_id = ?", workspaceID)
}
if err = memberQuery.Order("id asc").First(&member).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "工作空间数据缺失")
return
}
@@ -203,6 +214,39 @@ func (h *AuthHandler) Logout(c *gin.Context) {
response.OK(c, nil)
}
func (h *AuthHandler) ChangePassword(c *gin.Context) {
var req changePasswordReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "密码格式不合法")
return
}
var user models.User
if err := h.DB.First(&user, c.GetUint64("uid")).Error; err != nil {
response.Fail(c, http.StatusUnauthorized, 1002, "用户不存在")
return
}
ok, err := auth.VerifyPassword(req.CurrentPassword, user.PasswordHash)
if err != nil || !ok {
response.Fail(c, http.StatusBadRequest, 2002, "当前密码错误")
return
}
if req.CurrentPassword == req.NewPassword {
response.Fail(c, http.StatusBadRequest, 1001, "新密码不能与当前密码相同")
return
}
hash, err := auth.HashPassword(req.NewPassword)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "密码更新失败")
return
}
if err = h.DB.Model(&user).Update("password_hash", hash).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "密码更新失败")
return
}
response.Audit(c, "auth.password_change", "user:"+itoa(user.ID), nil)
response.OK(c, nil)
}
// Me 当前用户信息
func (h *AuthHandler) Me(c *gin.Context) {
var user models.User
@@ -213,6 +257,34 @@ func (h *AuthHandler) Me(c *gin.Context) {
response.OK(c, gin.H{"user": userDTO(user), "workspaceId": c.GetUint64("wsid"), "memberRole": c.GetString("mrole")})
}
// SwitchWorkspace issues a fresh access/refresh pair for another workspace
// where the current user is a member. The old token remains valid until its
// normal expiry, but all resource handlers use the newly selected workspace.
func (h *AuthHandler) SwitchWorkspace(c *gin.Context) {
wsid, 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("workspace_id = ? and user_id = ?", wsid, c.GetUint64("uid")).First(&member).Error; err != nil {
response.Fail(c, http.StatusForbidden, 1003, "不是该工作空间成员")
return
}
var user models.User
if err = h.DB.First(&user, c.GetUint64("uid")).Error; err != nil || user.Status != "active" {
response.Fail(c, http.StatusUnauthorized, 1002, "用户不可用")
return
}
access, refresh, err := h.issuePair(user.ID, wsid, user.Role, member.Role)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "签发令牌失败")
return
}
response.Audit(c, "workspace.switch", "workspace:"+itoa(wsid), nil)
response.OK(c, gin.H{"accessToken": access, "refreshToken": refresh, "workspace": gin.H{"id": wsid}, "memberRole": member.Role})
}
// 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)
@@ -220,7 +292,7 @@ func (h *AuthHandler) issuePair(uid, wsid uint64, role, mrole string) (string, s
return "", "", err
}
jti := uuid.NewString()
refresh, err := auth.IssueRefresh(h.Cfg.JWTSecret, uid, jti)
refresh, err := auth.IssueRefresh(h.Cfg.JWTSecret, uid, wsid, jti)
if err != nil {
return "", "", err
}
+95 -10
View File
@@ -3,6 +3,7 @@ package handlers
import (
"net/http"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
@@ -10,18 +11,35 @@ import (
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
"everypublish/server/internal/platform"
"everypublish/server/internal/ws"
)
// ChallengeHandler 二次验证挑战处理器
type ChallengeHandler struct {
DB *gorm.DB
DB *gorm.DB
Hub *ws.Hub
Registry *platform.Registry
ExecutorMode string
StartLoginPoll func(models.Challenge, models.Account, uint64, platform.Adapter)
}
// 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")
var expired []models.Challenge
h.DB.Where("workspace_id = ? and status = ? and expires_at < ?", c.GetUint64("wsid"), "active", time.Now()).Find(&expired)
if len(expired) > 0 {
ids := make([]uint64, 0, len(expired))
for _, challenge := range expired {
ids = append(ids, challenge.ID)
}
h.DB.Model(&models.Challenge{}).Where("id in ?", ids).Update("status", "expired")
accountIDs := make([]uint64, 0, len(expired))
for _, challenge := range expired {
accountIDs = append(accountIDs, challenge.AccountID)
}
_ = h.DB.Model(&models.Account{}).Where("workspace_id = ? and id in ? and status = ?", c.GetUint64("wsid"), accountIDs, "binding").Updates(map[string]interface{}{"status": "expired", "last_error": "登录二维码已过期"}).Error
}
q := h.DB.Model(&models.Challenge{}).Where("workspace_id = ?", c.GetUint64("wsid"))
if s := c.Query("status"); s != "" {
q = q.Where("status = ?", s)
@@ -34,7 +52,8 @@ func (h *ChallengeHandler) List(c *gin.Context) {
response.OK(c, gin.H{"list": list})
}
// Solve 人工完成挑战(验证码/APP确认;扫码由 Agent 完成)
// Solve 完成网页挑战。mock QR 只接受显式联调值;真实平台挑战必须由 adapter
// 在用户完成扫码后回写,不能通过空请求把账号标成 active。
func (h *ChallengeHandler) Solve(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
@@ -44,7 +63,10 @@ func (h *ChallengeHandler) Solve(c *gin.Context) {
var req struct {
Value string `json:"value"`
}
_ = c.ShouldBindJSON(&req)
if err = c.ShouldBindJSON(&req); 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, "挑战不存在")
@@ -54,16 +76,35 @@ func (h *ChallengeHandler) Solve(c *gin.Context) {
response.Fail(c, http.StatusConflict, 3003, "挑战已结束")
return
}
value := strings.TrimSpace(req.Value)
if value == "" {
response.Fail(c, http.StatusBadRequest, 1001, "挑战确认值不能为空")
return
}
if challenge.Kind == "qr" && !strings.HasPrefix(challenge.QRURL, "mock://") {
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
}
_ = h.DB.Model(&models.Account{}).
Where("id = ? and workspace_id = ?", challenge.AccountID, c.GetUint64("wsid")).
Updates(map[string]interface{}{"status": "active", "last_active_at": &now}).Error
_ = Notify(h.DB, c.GetUint64("wsid"), []uint64{c.GetUint64("uid")}, "challenge", "登录挑战已完成", "账号已恢复可用")
if h.Hub != nil {
h.Hub.BroadcastWS(c.GetUint64("wsid"), "challenge.status", gin.H{
"challengeId": challenge.ID, "accountId": challenge.AccountID, "status": "solved",
})
h.Hub.BroadcastWS(c.GetUint64("wsid"), "notification.new", gin.H{"kind": "challenge", "title": "登录挑战已完成", "content": "账号已恢复可用"})
}
response.Audit(c, "challenge.solve", "challenge:"+itoa(id), gin.H{"kind": challenge.Kind})
response.OK(c, nil)
}
// Resend 一键重发(重置过期时间,实际重发出 Agent 执行)
// Resend 一键重发(重置过期时间,交由当前 Web adapter 继续处理)
func (h *ChallengeHandler) Resend(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
@@ -79,14 +120,57 @@ func (h *ChallengeHandler) Resend(c *gin.Context) {
response.Fail(c, http.StatusConflict, 3003, "当前状态不可重发")
return
}
if err = h.DB.Model(&challenge).Updates(map[string]interface{}{
var account models.Account
if err = h.DB.Where("id = ? and workspace_id = ?", challenge.AccountID, c.GetUint64("wsid")).First(&account).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "账号不存在")
return
}
var adapter platform.Adapter
var session platform.LoginSession
if h.ExecutorMode != "mock" && h.Registry != nil {
adapter = h.Registry.Get(account.Platform)
if adapter != nil {
session, err = adapter.BeginLogin(c.Request.Context())
if err != nil {
response.Fail(c, http.StatusBadGateway, 2007, adapterErrorMessage(err))
return
}
}
}
updates := map[string]interface{}{
"status": "active", "expires_at": time.Now().Add(30 * time.Minute),
}).Error; err != nil {
"prompt": "请在网页中完成 " + account.Platform + " 登录",
}
if adapter != nil {
updates["qr_token"] = session.Token
updates["qr_url"] = session.URL
}
if err = h.DB.Model(&challenge).Updates(updates).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
_ = h.DB.Model(&account).Updates(map[string]interface{}{"status": "binding", "last_error": ""}).Error
if adapter != nil && h.StartLoginPoll != nil {
challenge.QRToken = session.Token
challenge.QRURL = session.URL
challenge.Status = "active"
challenge.ExpiresAt = updates["expires_at"].(time.Time)
challenge.Prompt = updates["prompt"].(string)
h.StartLoginPoll(challenge, account, c.GetUint64("uid"), adapter)
} else {
challenge.Status = "active"
challenge.ExpiresAt = updates["expires_at"].(time.Time)
challenge.Prompt = updates["prompt"].(string)
}
if h.Hub != nil {
h.Hub.BroadcastWS(challenge.WorkspaceID, "challenge.status", gin.H{
"challengeId": challenge.ID, "accountId": challenge.AccountID, "status": "active",
"kind": challenge.Kind, "platform": challenge.Platform, "qrUrl": challenge.QRURL,
"prompt": challenge.Prompt, "expiresAt": challenge.ExpiresAt,
})
}
response.Audit(c, "challenge.resend", "challenge:"+itoa(id), nil)
response.OK(c, nil)
response.OK(c, gin.H{"challengeId": challenge.ID, "qrUrl": challenge.QRURL, "qrToken": challenge.QRToken, "prompt": challenge.Prompt, "expiresAt": challenge.ExpiresAt})
}
// Suspend 挂起(超时/人工)
@@ -109,6 +193,7 @@ func (h *ChallengeHandler) Suspend(c *gin.Context) {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
_ = h.DB.Model(&models.Account{}).Where("id = ? and workspace_id = ? and status = ?", challenge.AccountID, c.GetUint64("wsid"), "binding").Updates(map[string]interface{}{"status": "expired", "last_error": "登录挑战已挂起"}).Error
response.Audit(c, "challenge.suspend", "challenge:"+itoa(id), nil)
response.OK(c, nil)
}
+17 -2
View File
@@ -95,7 +95,8 @@ func (h *MaterialHandler) Upload(c *gin.Context) {
return
}
kind := "image"
if strings.HasPrefix(header.Header.Get("Content-Type"), "video/") {
contentType := strings.ToLower(header.Header.Get("Content-Type"))
if strings.HasPrefix(contentType, "video/") || isVideoExtension(ext) {
kind = "video"
}
group := strings.TrimSpace(c.PostForm("group"))
@@ -125,6 +126,15 @@ func (h *MaterialHandler) Upload(c *gin.Context) {
response.OK(c, gin.H{"material": material, "dedup": false})
}
func isVideoExtension(ext string) bool {
switch strings.ToLower(ext) {
case ".mp4", ".mov", ".avi", ".mkv", ".m4v", ".webm", ".flv", ".wmv":
return true
default:
return false
}
}
// URL 生成签名直链(10 分钟一次性)
func (h *MaterialHandler) URL(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
@@ -171,10 +181,15 @@ func (h *MaterialHandler) ServeFile(c *gin.Context) {
_ = 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)) {
root := filepath.Clean(h.Cfg.StorageDir) + string(os.PathSeparator)
if !strings.HasPrefix(clean, root) {
response.Fail(c, http.StatusForbidden, 1003, "非法路径")
return
}
if c.Query("inline") == "1" {
c.File(abspath)
return
}
c.FileAttachment(abspath, material.Name)
}
+1 -1
View File
@@ -153,7 +153,7 @@ func (h *MemberHandler) UpdateRole(c *gin.Context) {
response.Fail(c, http.StatusForbidden, 1003, "不可修改 owner 角色")
return
}
if err = h.DB.Model(&member).Update("role", req.Role).Error; err != nil {
if err = h.DB.Model(&member).Where("workspace_id = ?", c.GetUint64("wsid")).Update("role", req.Role).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
+104 -2
View File
@@ -2,6 +2,7 @@ package handlers
import (
"encoding/json"
"fmt"
"net/http"
"strconv"
"time"
@@ -39,6 +40,54 @@ func marshalJSONArray[T any](arr []T) string {
return string(raw)
}
func (h *TaskHandler) validateTargets(wsid uint64, accountIDs, materialIDs []uint64) error {
var accounts []models.Account
if err := h.DB.Where("workspace_id = ? and id in ? and status = ?", wsid, accountIDs, "active").Find(&accounts).Error; err != nil {
return err
}
if len(accounts) != len(uniqueIDs(accountIDs)) {
return fmt.Errorf("账号不存在、不属于当前工作区或尚未完成登录")
}
var materials []models.Material
if err := h.DB.Where("workspace_id = ? and id in ? and status = ?", wsid, materialIDs, "ready").Find(&materials).Error; err != nil {
return err
}
if len(materials) != len(uniqueIDs(materialIDs)) {
return fmt.Errorf("素材不存在、未就绪或不属于当前工作区")
}
return nil
}
func uniqueIDs(values []uint64) []uint64 {
seen := make(map[uint64]struct{}, len(values))
result := make([]uint64, 0, len(values))
for _, value := range values {
if _, ok := seen[value]; ok {
continue
}
seen[value] = struct{}{}
result = append(result, value)
}
return result
}
func recordTaskEvent(db *gorm.DB, row *models.Task, eventType, message string, data interface{}) {
raw := ""
if data != nil {
if encoded, err := json.Marshal(data); err == nil {
raw = string(encoded)
}
}
_ = db.Create(&models.TaskEvent{
WorkspaceID: row.WorkspaceID,
TaskID: row.ID,
Type: eventType,
Status: row.Status,
Message: message,
Data: raw,
}).Error
}
// List 任务列表(状态筛选 + 分页)
func (h *TaskHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
@@ -46,7 +95,7 @@ func (h *TaskHandler) List(c *gin.Context) {
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
if size < 1 || size > 300 {
size = 20
}
q := h.DB.Model(&models.Task{}).Where("workspace_id = ?", c.GetUint64("wsid"))
@@ -82,6 +131,10 @@ func (h *TaskHandler) Create(c *gin.Context) {
response.Fail(c, http.StatusBadRequest, 1001, "至少选择一个素材")
return
}
if err := h.validateTargets(c.GetUint64("wsid"), req.AccountIDs, req.MaterialIDs); err != nil {
response.Fail(c, http.StatusBadRequest, 1004, err.Error())
return
}
priority := req.Priority
if priority < 1 || priority > 10 {
priority = 5
@@ -100,16 +153,36 @@ func (h *TaskHandler) Create(c *gin.Context) {
}
if req.ScheduleAt != nil {
ts := time.UnixMilli(*req.ScheduleAt)
if !ts.After(time.Now()) {
response.Fail(c, http.StatusBadRequest, 1001, "定时发布时间必须晚于当前时间")
return
}
row.ScheduleAt = &ts
}
if err := h.DB.Create(&row).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "创建失败")
return
}
recordTaskEvent(h.DB, &row, "created", "任务已创建", nil)
response.Audit(c, "task.create", "task:"+itoa(row.ID), gin.H{"title": req.Title})
response.OK(c, row)
}
// Events 返回任务时间线,并强制按 workspace 隔离。
func (h *TaskHandler) Events(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var events []models.TaskEvent
if err = h.DB.Where("task_id = ? and workspace_id = ?", id, c.GetUint64("wsid")).Order("id asc").Find(&events).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "任务时间线加载失败")
return
}
response.OK(c, gin.H{"list": events})
}
// Get 任务详情
func (h *TaskHandler) Get(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
@@ -146,6 +219,14 @@ func (h *TaskHandler) Update(c *gin.Context) {
response.Fail(c, http.StatusConflict, 3005, "仅草稿或已驳回任务可编辑")
return
}
if len(req.AccountIDs) == 0 || len(req.MaterialIDs) == 0 {
response.Fail(c, http.StatusBadRequest, 1001, "至少选择一个账号和素材")
return
}
if err := h.validateTargets(c.GetUint64("wsid"), req.AccountIDs, req.MaterialIDs); err != nil {
response.Fail(c, http.StatusBadRequest, 1004, err.Error())
return
}
updates := map[string]interface{}{
"title": req.Title,
"content": req.Content,
@@ -155,7 +236,15 @@ func (h *TaskHandler) Update(c *gin.Context) {
}
if req.ScheduleAt != nil {
ts := time.UnixMilli(*req.ScheduleAt)
if !ts.After(time.Now()) {
response.Fail(c, http.StatusBadRequest, 1001, "定时发布时间必须晚于当前时间")
return
}
updates["schedule_at"] = &ts
} else {
// nil represents an explicit immediate schedule in the JSON contract.
// The update endpoint has no meaningful partial-update semantics.
updates["schedule_at"] = nil
}
if err = h.DB.Model(&row).Updates(updates).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
@@ -215,6 +304,7 @@ func (h *TaskHandler) apply(c *gin.Context, action task.Action, extra map[string
return nil
}
row.Status = next
recordTaskEvent(h.DB, &row, string(action), "状态变更为 "+next, extra)
return &row
}
@@ -230,9 +320,15 @@ func (h *TaskHandler) Submit(c *gin.Context) {
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)
if h.Dispatcher != nil && h.Dispatcher.Hub != nil {
h.Dispatcher.Hub.BroadcastWS(row.WorkspaceID, "notification.new", gin.H{"kind": "task", "title": "任务已通过审核", "content": row.Title})
}
response.Audit(c, "task.approve", "task:"+itoa(row.ID), nil)
if h.Dispatcher != nil {
_ = h.Dispatcher.TryDispatch(row.ID) // 在线设备立即下发
_ = h.Dispatcher.TryDispatch(row.ID)
// TryDispatch can synchronously advance a Web-only mock task. Return
// the persisted state instead of the stale queued copy from apply.
_ = h.DB.Where("id = ? and workspace_id = ?", row.ID, row.WorkspaceID).First(row).Error
}
response.OK(c, row)
}
@@ -246,6 +342,9 @@ func (h *TaskHandler) Reject(c *gin.Context) {
_ = 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)
if h.Dispatcher != nil && h.Dispatcher.Hub != nil {
h.Dispatcher.Hub.BroadcastWS(row.WorkspaceID, "notification.new", gin.H{"kind": "task", "title": "任务被驳回", "content": req.Note})
}
response.Audit(c, "task.reject", "task:"+itoa(row.ID), gin.H{"note": req.Note})
response.OK(c, row)
}
@@ -273,6 +372,9 @@ func (h *TaskHandler) Retry(c *gin.Context) {
// Cancel 取消
func (h *TaskHandler) Cancel(c *gin.Context) {
if row := h.apply(c, task.ActionCancel, nil); row != nil {
if h.Dispatcher != nil {
h.Dispatcher.Cancel(row.ID)
}
response.Audit(c, "task.cancel", "task:"+itoa(row.ID), nil)
response.OK(c, row)
}
@@ -75,6 +75,10 @@ func (h *WorkspaceHandler) Update(c *gin.Context) {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
if id != c.GetUint64("wsid") {
response.Fail(c, http.StatusNotFound, 1004, "工作空间不存在")
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
+27
View File
@@ -0,0 +1,27 @@
package middleware
import (
"net/http"
"github.com/gin-gonic/gin"
"everypublish/server/internal/api/response"
)
// WorkspaceRoles enforces the member role from the current JWT. Resource
// handlers still apply workspace_id predicates; this middleware only handles
// the action-level permission boundary.
func WorkspaceRoles(roles ...string) gin.HandlerFunc {
allowed := make(map[string]struct{}, len(roles))
for _, role := range roles {
allowed[role] = struct{}{}
}
return func(c *gin.Context) {
if _, ok := allowed[c.GetString("mrole")]; !ok {
response.Fail(c, http.StatusForbidden, 1003, "当前角色无权执行此操作")
c.Abort()
return
}
c.Next()
}
}
+50 -35
View File
@@ -1,7 +1,6 @@
package api
import (
"crypto/ed25519"
"net/http"
"time"
@@ -12,11 +11,13 @@ import (
"everypublish/server/internal/api/middleware"
"everypublish/server/internal/cache"
"everypublish/server/internal/config"
"everypublish/server/internal/credentials"
"everypublish/server/internal/platform"
"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 {
func Router(cfg *config.Config, db *gorm.DB, rds *cache.Redis, hub *ws.Hub, dsp *ws.Dispatcher) *gin.Engine {
r := gin.New()
r.Use(gin.Logger(), gin.Recovery())
@@ -35,23 +36,41 @@ func Router(cfg *config.Config, db *gorm.DB, rds *cache.Redis, hub *ws.Hub, serv
})
})
// 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))
api.GET("/runtime/status", middleware.AuthRequired(cfg), func(c *gin.Context) {
mode := cfg.ExecutorMode
adapters := []string(nil)
if dsp != nil && dsp.Registry != nil {
adapters = dsp.Registry.Names()
}
dbOK := false
if sqlDB, err := db.DB(); err == nil {
dbOK = sqlDB.Ping() == nil
}
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": gin.H{
"executorMode": mode, "adapters": adapters, "db": dbOK, "redis": rds.Ping() == nil,
}})
})
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}
var registry *platform.Registry
var credentialStore *credentials.Store
if dsp != nil {
registry = dsp.Registry
credentialStore = dsp.Credentials
}
accH := &handlers.AccountHandler{DB: db, Hub: hub, Registry: registry, Credentials: credentialStore, ExecutorMode: cfg.ExecutorMode}
matH := &handlers.MaterialHandler{DB: db, Cfg: cfg}
tskH := &handlers.TaskHandler{DB: db, Dispatcher: dsp}
chlH := &handlers.ChallengeHandler{DB: db}
chlH := &handlers.ChallengeHandler{DB: db, Hub: hub, Registry: registry, ExecutorMode: cfg.ExecutorMode}
chlH.StartLoginPoll = accH.StartLoginPoll
ntfH := &handlers.NotificationHandler{DB: db}
agH := &handlers.AgentHandler{DB: db, Hub: hub}
auth := api.Group("/auth")
{
@@ -59,7 +78,9 @@ func Router(cfg *config.Config, db *gorm.DB, rds *cache.Redis, hub *ws.Hub, serv
auth.POST("/login", authH.Login)
auth.POST("/refresh", authH.Refresh)
auth.POST("/logout", middleware.AuthRequired(cfg), authH.Logout)
auth.POST("/password", middleware.AuthRequired(cfg), authH.ChangePassword)
auth.GET("/me", middleware.AuthRequired(cfg), authH.Me)
auth.POST("/switch-workspace/:id", middleware.AuthRequired(cfg), authH.SwitchWorkspace)
}
wsg := api.Group("/workspaces", middleware.AuthRequired(cfg))
@@ -86,17 +107,20 @@ func Router(cfg *config.Config, db *gorm.DB, rds *cache.Redis, hub *ws.Hub, serv
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)
acc.POST("", middleware.WorkspaceRoles("owner", "admin", "operator"), accH.Create)
acc.PUT("/:id", middleware.WorkspaceRoles("owner", "admin", "operator"), accH.Update)
acc.DELETE("/:id", middleware.WorkspaceRoles("owner", "admin"), accH.Delete)
acc.POST("/:id/bind", middleware.WorkspaceRoles("owner", "admin", "operator"), accH.Bind)
acc.POST("/:id/credentials", middleware.WorkspaceRoles("owner", "admin", "operator"), accH.ImportCredentials)
acc.POST("/:id/check", middleware.WorkspaceRoles("owner", "admin", "operator"), accH.Check)
acc.POST("/:id/unbind", middleware.WorkspaceRoles("owner", "admin"), accH.Unbind)
}
mat := api.Group("/materials", middleware.AuthRequired(cfg))
{
mat.GET("", matH.List)
mat.POST("", matH.Upload)
mat.DELETE("/:id", matH.Delete)
mat.POST("", middleware.WorkspaceRoles("owner", "admin", "operator"), matH.Upload)
mat.DELETE("/:id", middleware.WorkspaceRoles("owner", "admin", "operator"), matH.Delete)
mat.GET("/:id/url", matH.URL)
}
@@ -108,24 +132,25 @@ func Router(cfg *config.Config, db *gorm.DB, rds *cache.Redis, hub *ws.Hub, serv
tsk := api.Group("/tasks", middleware.AuthRequired(cfg))
{
tsk.GET("", tskH.List)
tsk.POST("", tskH.Create)
tsk.POST("", middleware.WorkspaceRoles("owner", "admin", "operator"), tskH.Create)
tsk.GET("/:id/events", tskH.Events)
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)
tsk.PUT("/:id", middleware.WorkspaceRoles("owner", "admin", "operator"), tskH.Update)
tsk.DELETE("/:id", middleware.WorkspaceRoles("owner", "admin", "operator"), tskH.Delete)
tsk.POST("/:id/submit", middleware.WorkspaceRoles("owner", "admin", "operator"), tskH.Submit)
tsk.POST("/:id/approve", middleware.WorkspaceRoles("owner", "admin", "reviewer"), tskH.Approve)
tsk.POST("/:id/reject", middleware.WorkspaceRoles("owner", "admin", "reviewer"), tskH.Reject)
tsk.POST("/:id/resubmit", middleware.WorkspaceRoles("owner", "admin", "operator"), tskH.Resubmit)
tsk.POST("/:id/retry", middleware.WorkspaceRoles("owner", "admin", "operator"), tskH.Retry)
tsk.POST("/:id/cancel", middleware.WorkspaceRoles("owner", "admin", "operator"), 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)
chl.POST("/:id/solve", middleware.WorkspaceRoles("owner", "admin", "operator", "reviewer"), chlH.Solve)
chl.POST("/:id/resend", middleware.WorkspaceRoles("owner", "admin", "operator", "reviewer"), chlH.Resend)
chl.POST("/:id/suspend", middleware.WorkspaceRoles("owner", "admin", "operator", "reviewer"), chlH.Suspend)
}
ntf := api.Group("/notifications", middleware.AuthRequired(cfg))
@@ -135,21 +160,11 @@ func Router(cfg *config.Config, db *gorm.DB, rds *cache.Redis, hub *ws.Hub, serv
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
+61 -17
View File
@@ -2,8 +2,6 @@ package api_test
import (
"bytes"
"crypto/ed25519"
"crypto/rand"
"encoding/json"
"fmt"
"net/http/httptest"
@@ -36,33 +34,33 @@ func TestMain(m *testing.M) {
adm, err := gorm.Open(gormmysql.Open(adminDSN), &gorm.Config{})
if err != nil {
fmt.Println("skip: mysql admin connect failed:", err)
os.Exit(0)
os.Exit(1)
}
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)
os.Exit(1)
}
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)
os.Exit(1)
}
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",
ServerAddr: ":0",
MySQLDSN: testDSN,
RedisAddr: "127.0.0.1:6379",
JWTSecret: "test-secret",
BaseURL: "http://127.0.0.1:8090",
StorageDir: "/tmp/everypublish-test",
ExecutorMode: "mock",
}
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)
testRouter = api.Router(cfg, testDB, rds, testHub, dsp)
code := m.Run()
os.Exit(code)
}
@@ -73,8 +71,8 @@ func resetTestDB() {
testHub.Clear()
}
for _, t := range []interface{}{
&models.TransferToken{}, &models.PairingCode{}, &models.Notification{}, &models.AuditLog{},
&models.Challenge{}, &models.AgentDevice{}, &models.Task{}, &models.Material{},
&models.TransferToken{}, &models.Notification{}, &models.AuditLog{},
&models.Challenge{}, &models.AgentDevice{}, &models.PairingCode{}, &models.TaskEvent{}, &models.Task{}, &models.Material{},
&models.Account{}, &models.Member{}, &models.Workspace{}, &models.User{},
} {
testDB.Migrator().DropTable(t)
@@ -174,6 +172,24 @@ func TestWorkspaceAndMemberFlow(t *testing.T) {
if code != 200 || env.Code != 0 {
t.Fatalf("create ws failed: %d %s", code, env.Message)
}
var secondWS models.Workspace
_ = json.Unmarshal(env.Data, &secondWS)
code, _ = doReq(t, "PUT", fmt.Sprintf("/api/v1/workspaces/%d", secondWS.ID), aAccess, gin.H{"name": "越权空间"})
if code != 404 {
t.Fatalf("cross-workspace update should be 404, got %d", code)
}
var ownWorkspaces struct {
List []models.Workspace `json:"list"`
}
code, env = doReq(t, "GET", "/api/v1/workspaces", aAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("list own workspaces failed: %d %s", code, env.Message)
}
_ = json.Unmarshal(env.Data, &ownWorkspaces)
if len(ownWorkspaces.List) == 0 {
t.Fatal("owner should have at least one workspace")
}
primaryWSID := ownWorkspaces.List[0].ID
// 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 {
@@ -192,6 +208,17 @@ func TestWorkspaceAndMemberFlow(t *testing.T) {
if code != 200 || env.Code != 0 {
t.Fatalf("join failed: %d %s", code, env.Message)
}
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/auth/switch-workspace/%d", primaryWSID), bAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("workspace switch failed: %d %s", code, env.Message)
}
var switched tokens
_ = json.Unmarshal(env.Data, &switched)
bAccess = switched.AccessToken
code, _ = doReq(t, "POST", "/api/v1/accounts", bAccess, gin.H{"platform": "douyin", "remark": "viewer-forbidden"})
if code != 403 {
t.Fatalf("viewer account write should be 403, got %d", code)
}
// A 查成员:2 人
code, env = doReq(t, "GET", "/api/v1/members", aAccess, nil)
if code != 200 || env.Code != 0 {
@@ -219,6 +246,23 @@ func TestWorkspaceAndMemberFlow(t *testing.T) {
if code != 200 || env.Code != 0 {
t.Fatalf("update role failed: %d %s", code, env.Message)
}
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/auth/switch-workspace/%d", primaryWSID), bAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("workspace switch after role update failed: %d %s", code, env.Message)
}
_ = json.Unmarshal(env.Data, &switched)
var switchInfo struct {
MemberRole string `json:"memberRole"`
}
_ = json.Unmarshal(env.Data, &switchInfo)
if switchInfo.MemberRole != "operator" {
t.Fatalf("workspace switch should reflect updated role, got %q", switchInfo.MemberRole)
}
bAccess = switched.AccessToken
code, env = doReq(t, "POST", "/api/v1/accounts", bAccess, gin.H{"platform": "douyin", "remark": "operator-allowed"})
if code != 200 || env.Code != 0 {
t.Fatalf("operator account write should be allowed: %d %s", code, env.Message)
}
// owner 不可被移除
code, _ = doReq(t, "DELETE", fmt.Sprintf("/api/v1/members/%d", ownerMember), aAccess, nil)
if code != 403 {
@@ -240,8 +284,8 @@ func TestWorkspaceAndMemberFlow(t *testing.T) {
t.Fatalf("audit list failed: %d %s", code, env.Message)
}
var al struct {
List []models.AuditLog `json:"list"`
Total int64 `json:"total"`
List []models.AuditLog `json:"list"`
Total int64 `json:"total"`
}
_ = json.Unmarshal(env.Data, &al)
if al.Total < 3 {
+43 -224
View File
@@ -2,246 +2,65 @@ 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)
// TestBrowserWSOnly verifies the Web-only realtime contract. Agent/device WSS
// is intentionally not part of the V1 router.
func TestBrowserWSOnly(t *testing.T) {
resetTestDB()
access, _ := register(t, "browser-ws@test.com")
srv := httptest.NewServer(testRouter)
defer srv.Close()
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws/browser?token=" + access
conn, _, err := websocket.Dial(context.Background(), wsURL, nil)
if err != nil {
t.Fatalf("browser websocket dial failed: %v", err)
}
defer conn.Close(websocket.StatusNormalClosure, "done")
code, env := doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "douyin", "remark": "browser ws"})
if code != 200 || env.Code != 0 {
t.Fatalf("create account failed: %d %s", code, env.Message)
}
var account struct {
ID uint64 `json:"id"`
}
if err := json.Unmarshal(env.Data, &account); err != nil || account.ID == 0 {
t.Fatalf("account response missing id: %s", env.Data)
}
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/accounts/%d/bind", account.ID), access, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("bind failed: %d %s", code, env.Message)
}
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
_, raw, err := conn.Read(ctx)
if err != nil {
t.Fatalf("ws read: %v", err)
t.Fatalf("browser websocket did not receive challenge event: %v", err)
}
var env proto.Envelope
if err = json.Unmarshal(raw, &env); err != nil {
t.Fatalf("ws unmarshal: %v", err)
var event struct {
Type string `json:"type"`
}
if err := json.Unmarshal(raw, &event); err != nil {
t.Fatal(err)
}
if event.Type != "challenge.status" && event.Type != "notification.new" {
t.Fatalf("unexpected browser event type: %s", event.Type)
}
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)
code, _ = doReq(t, "GET", "/api/v1/agent/devices", access, nil)
if code != 404 {
t.Fatalf("agent route must be absent in Web-only mode, got %d", code)
}
}
+2 -2
View File
@@ -52,7 +52,7 @@ func TestJWTAccessExpired(t *testing.T) {
func TestJWTRefreshRoundTrip(t *testing.T) {
secret := "test-secret"
token, err := IssueRefresh(secret, 9, "jti-1")
token, err := IssueRefresh(secret, 9, 7, "jti-1")
if err != nil {
t.Fatalf("issue: %v", err)
}
@@ -60,7 +60,7 @@ func TestJWTRefreshRoundTrip(t *testing.T) {
if err != nil {
t.Fatalf("parse: %v", err)
}
if claims.UID != 9 || claims.JTI != "jti-1" {
if claims.UID != 9 || claims.WorkspaceID != 7 || claims.JTI != "jti-1" {
t.Fatalf("claims mismatch: %+v", claims)
}
}
+5 -4
View File
@@ -25,8 +25,9 @@ type AccessClaims struct {
// RefreshClaims 刷新令牌声明
type RefreshClaims struct {
UID uint64 `json:"uid"`
JTI string `json:"jti"`
UID uint64 `json:"uid"`
WorkspaceID uint64 `json:"wsid"`
JTI string `json:"jti"`
jwt.RegisteredClaims
}
@@ -53,10 +54,10 @@ func IssueAccess(secret string, uid, wsid uint64, role, memberRole string) (stri
}
// IssueRefresh 签发刷新令牌
func IssueRefresh(secret string, uid uint64, jti string) (string, error) {
func IssueRefresh(secret string, uid, wsid uint64, jti string) (string, error) {
now := time.Now()
claims := RefreshClaims{
UID: uid, JTI: jti,
UID: uid, WorkspaceID: wsid, JTI: jti,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(RefreshTTL)),
+30 -20
View File
@@ -7,16 +7,21 @@ import (
// 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 上传体积上限
ServerAddr string
MySQLDSN string
RedisAddr string
RedisPass string
JWTSecret string
BaseURL string
PublicBaseURL string // 浏览器侧访问地址(邀请链接等)
StorageDir string
CredentialDir string // 平台登录凭据加密目录
StaticDir string // 前端生产构建目录(空=不托管静态)
MaxUploadBytes int64 // multipart 上传体积上限
ExecutorMode string // mock=本地联调;web=启用已注册平台适配器并保留未接入平台 mock fallback
BilibiliMemberBase string
BilibiliPassportBase string
BilibiliUpOSScheme string
}
// Load 从环境变量加载
@@ -28,16 +33,21 @@ func Load() *Config {
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),
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"),
CredentialDir: getenv("CREDENTIAL_DIR", "./data/credentials"),
StaticDir: getenv("STATIC_DIR", "../apps/web/dist"),
MaxUploadBytes: getenvInt64("MAX_UPLOAD_BYTES", 2<<30),
ExecutorMode: getenv("EXECUTOR_MODE", "mock"),
BilibiliMemberBase: getenv("BILIBILI_MEMBER_BASE", "https://member.bilibili.com"),
BilibiliPassportBase: getenv("BILIBILI_PASSPORT_BASE", "https://passport.bilibili.com"),
BilibiliUpOSScheme: getenv("BILIBILI_UPOS_SCHEME", "https"),
}
}
+193
View File
@@ -0,0 +1,193 @@
package credentials
// Package credentials stores platform secrets outside the SQL account model.
// Files are encrypted with AES-256-GCM and keyed by workspace/account IDs so a
// leaked account row or browser event never contains cookies.
import (
"crypto/aes"
"crypto/cipher"
"crypto/rand"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strconv"
"sync"
)
var errInvalidID = errors.New("workspace and account IDs must be positive")
// Store is a process-safe encrypted credential store.
type Store struct {
root string
key []byte
mu sync.Mutex
}
// Open creates or loads a 32-byte master key in root. The key file is never
// returned by an API and is permission-restricted to the current user.
func Open(root string) (*Store, error) {
if root == "" {
return nil, errors.New("credential directory is empty")
}
if err := os.MkdirAll(root, 0o700); err != nil {
return nil, err
}
keyPath := filepath.Join(root, "master.key")
key, err := os.ReadFile(keyPath)
if errors.Is(err, os.ErrNotExist) {
key = make([]byte, 32)
if _, err = io.ReadFull(rand.Reader, key); err != nil {
return nil, err
}
if err = writePrivate(keyPath, key); err != nil {
return nil, err
}
} else if err != nil {
return nil, err
}
if len(key) != 32 {
return nil, errors.New("credential master key must be 32 bytes")
}
return &Store{root: root, key: append([]byte(nil), key...)}, nil
}
func (s *Store) path(workspaceID, accountID uint64) (string, error) {
if workspaceID == 0 || accountID == 0 {
return "", errInvalidID
}
dir := filepath.Join(s.root, strconv.FormatUint(workspaceID, 10))
return filepath.Join(dir, strconv.FormatUint(accountID, 10)+".enc"), nil
}
// Save encrypts value and atomically replaces the account file.
func (s *Store) Save(workspaceID, accountID uint64, value []byte) error {
path, err := s.path(workspaceID, accountID)
if err != nil {
return err
}
block, err := aes.NewCipher(s.key)
if err != nil {
return err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return err
}
nonce := make([]byte, gcm.NonceSize())
if _, err = io.ReadFull(rand.Reader, nonce); err != nil {
return err
}
ciphertext := gcm.Seal(nil, nonce, value, nil)
payload := struct {
Nonce string `json:"nonce"`
Data string `json:"data"`
}{base64.RawStdEncoding.EncodeToString(nonce), base64.RawStdEncoding.EncodeToString(ciphertext)}
raw, err := json.Marshal(payload)
if err != nil {
return err
}
s.mu.Lock()
defer s.mu.Unlock()
if err = os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
return err
}
tmp, err := os.CreateTemp(filepath.Dir(path), ".credential-*")
if err != nil {
return err
}
tmpName := tmp.Name()
defer os.Remove(tmpName)
if err = tmp.Chmod(0o600); err == nil {
_, err = tmp.Write(raw)
}
if closeErr := tmp.Close(); err == nil {
err = closeErr
}
if err != nil {
return err
}
return os.Rename(tmpName, path)
}
// Load decrypts credentials for one workspace/account pair.
func (s *Store) Load(workspaceID, accountID uint64) ([]byte, error) {
path, err := s.path(workspaceID, accountID)
if err != nil {
return nil, err
}
s.mu.Lock()
raw, err := os.ReadFile(path)
s.mu.Unlock()
if err != nil {
return nil, err
}
var payload struct {
Nonce string `json:"nonce"`
Data string `json:"data"`
}
if err = json.Unmarshal(raw, &payload); err != nil {
return nil, fmt.Errorf("credential envelope: %w", err)
}
nonce, err := base64.RawStdEncoding.DecodeString(payload.Nonce)
if err != nil {
return nil, fmt.Errorf("credential nonce: %w", err)
}
ciphertext, err := base64.RawStdEncoding.DecodeString(payload.Data)
if err != nil {
return nil, fmt.Errorf("credential data: %w", err)
}
block, err := aes.NewCipher(s.key)
if err != nil {
return nil, err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return nil, err
}
if len(nonce) != gcm.NonceSize() {
return nil, errors.New("credential nonce has invalid size")
}
plain, err := gcm.Open(nil, nonce, ciphertext, nil)
if err != nil {
return nil, errors.New("credential authentication failed")
}
return plain, nil
}
// Delete removes credentials. Missing files are treated as success.
func (s *Store) Delete(workspaceID, accountID uint64) error {
path, err := s.path(workspaceID, accountID)
if err != nil {
return err
}
s.mu.Lock()
defer s.mu.Unlock()
if err = os.Remove(path); errors.Is(err, os.ErrNotExist) {
return nil
}
return err
}
func writePrivate(path string, value []byte) error {
tmp, err := os.CreateTemp(filepath.Dir(path), ".master-*")
if err != nil {
return err
}
tmpName := tmp.Name()
defer os.Remove(tmpName)
if err = tmp.Chmod(0o600); err == nil {
_, err = tmp.Write(value)
}
if closeErr := tmp.Close(); err == nil {
err = closeErr
}
if err != nil {
return err
}
return os.Rename(tmpName, path)
}
+49
View File
@@ -0,0 +1,49 @@
package credentials
import (
"os"
"path/filepath"
"testing"
)
func TestStoreRoundTripAndIsolation(t *testing.T) {
root := t.TempDir()
s, err := Open(root)
if err != nil {
t.Fatal(err)
}
secret := []byte(`{"cookies":"SESSDATA=secret"}`)
if err = s.Save(7, 11, secret); err != nil {
t.Fatal(err)
}
got, err := s.Load(7, 11)
if err != nil {
t.Fatal(err)
}
if string(got) != string(secret) {
t.Fatalf("round trip mismatch: %q", got)
}
if _, err = s.Load(8, 11); !os.IsNotExist(err) {
t.Fatalf("workspace isolation failed: %v", err)
}
raw, err := os.ReadFile(filepath.Join(root, "7", "11.enc"))
if err != nil {
t.Fatal(err)
}
if string(raw) == string(secret) {
t.Fatal("credential file is plaintext")
}
if err = s.Delete(7, 11); err != nil {
t.Fatal(err)
}
}
func TestStoreRejectsInvalidIDs(t *testing.T) {
s, err := Open(t.TempDir())
if err != nil {
t.Fatal(err)
}
if err = s.Save(0, 1, []byte("x")); err == nil {
t.Fatal("expected invalid ID error")
}
}
+1
View File
@@ -43,6 +43,7 @@ func AutoMigrate(db *gorm.DB) error {
&models.Account{},
&models.Material{},
&models.Task{},
&models.TaskEvent{},
&models.AgentDevice{},
&models.Challenge{},
&models.AuditLog{},
+23 -12
View File
@@ -38,20 +38,20 @@ type Member struct {
Role string `gorm:"size:16;default:viewer" json:"role"`
}
// Account 平台账号(凭据仅在 Agent 本机保险库,服务器零凭据)
// Account 平台账号(Web-only V1 的本机执行器持有登录 profile;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"`
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:"-"` // legacy column; not part of Web-only API
IPProfile string `gorm:"size:191" json:"ipProfile"`
LastActiveAt *time.Time `json:"lastActiveAt"`
LastError string `gorm:"size:512" json:"lastError"`
}
// Material 素材
@@ -93,6 +93,17 @@ type Task struct {
ReviewNote string `gorm:"size:512" json:"reviewNote"`
}
// TaskEvent 任务状态时间线(用于网页详情、审计和失败定位)。
type TaskEvent struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
TaskID uint64 `gorm:"index" json:"taskId"`
Type string `gorm:"size:32" json:"type"`
Status string `gorm:"size:32" json:"status"`
Message string `gorm:"size:1024" json:"message"`
Data string `gorm:"type:text" json:"data"`
}
// AgentDevice 客户端设备(配对时注册 Ed25519 公钥)
type AgentDevice struct {
Base
@@ -0,0 +1,550 @@
package bilibili
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"path/filepath"
"strconv"
"strings"
"time"
"everypublish/server/internal/credentials"
"everypublish/server/internal/platform"
)
// Adapter implements the server-side Bilibili API boundary. The exported
// endpoint fields make the adapter straightforward to use with an httptest
// server while retaining production defaults in New.
type Adapter struct {
Store *credentials.Store
Credentials *credentials.Store // compatibility alias for callers using this name
HTTP *http.Client
MemberBase string
PassportBase string
UpOSScheme string
PollTimeout time.Duration
PollInterval time.Duration
MaxRetries int
}
func New(store *credentials.Store) *Adapter {
return &Adapter{Store: store, Credentials: store, HTTP: &http.Client{Timeout: 30 * time.Second},
MemberBase: "https://member.bilibili.com", PassportBase: "https://passport.bilibili.com",
UpOSScheme: "https", PollTimeout: 3 * time.Minute, PollInterval: 2 * time.Second, MaxRetries: 2}
}
func (a *Adapter) Platform() string { return "bilibili" }
func (a *Adapter) BeginLogin(ctx context.Context) (platform.LoginSession, error) {
var ret struct {
Code int `json:"code"`
Data struct {
URL string `json:"url"`
Key string `json:"qrcode_key"`
} `json:"data"`
}
if err := a.getJSON(ctx, a.passportBase()+"/x/passport-login/web/qrcode/generate", "", &ret); err != nil {
return platform.LoginSession{}, err
}
if ret.Code != 0 || ret.Data.Key == "" || ret.Data.URL == "" {
return platform.LoginSession{}, perr(platform.PlatformChanged, "invalid QR response", nil)
}
return platform.LoginSession{Token: ret.Data.Key, URL: ret.Data.URL}, nil
}
func (a *Adapter) PollLogin(ctx context.Context, token string) (platform.LoginPoll, error) {
if token == "" {
return platform.LoginPoll{}, perr(platform.Validation, "QR token is empty", nil)
}
q := url.Values{"qrcode_key": []string{token}}
var ret struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
URL string `json:"url"`
} `json:"data"`
}
if err := a.getJSON(ctx, a.passportBase()+"/x/passport-login/web/qrcode/poll?"+q.Encode(), "", &ret); err != nil {
return platform.LoginPoll{}, err
}
switch ret.Code {
case 86101:
return platform.LoginPoll{State: platform.LoginPending, Prompt: "请使用哔哩哔哩手机客户端扫码"}, nil
case 86090:
return platform.LoginPoll{State: platform.LoginScanned, Prompt: "已扫码,请在哔哩哔哩手机客户端确认登录"}, nil
case 0:
cookies := cookiesFromURL(ret.Data.URL)
if cookies == "" {
return platform.LoginPoll{}, perr(platform.PlatformChanged, "login response did not contain cookies", nil)
}
return platform.LoginPoll{State: platform.LoginConfirmed, Cookies: cookies}, nil
default:
return platform.LoginPoll{}, perr(platform.PlatformChanged, "QR poll rejected: "+ret.Message, nil)
}
}
func (a *Adapter) SaveCredentials(workspaceID, accountID uint64, cookies string) error {
if strings.TrimSpace(cookies) == "" {
return perr(platform.Validation, "cookies are empty", nil)
}
if cookieValue(cookies, "SESSDATA") == "" || cookieValue(cookies, "bili_jct") == "" || cookieValue(cookies, "DedeUserID") == "" {
return perr(platform.Validation, "Bilibili cookies must include SESSDATA, bili_jct and DedeUserID", nil)
}
s := a.store()
if s == nil {
return perr(platform.Unknown, "credential store is not configured", nil)
}
if err := s.Save(workspaceID, accountID, []byte(cookies)); err != nil {
return perr(platform.Unknown, "saving credentials", err)
}
return nil
}
func (a *Adapter) CheckCredentials(ctx context.Context, workspaceID, accountID uint64) error {
s := a.store()
if s == nil {
return perr(platform.Unknown, "credential store is not configured", nil)
}
raw, err := s.Load(workspaceID, accountID)
if err != nil {
return perr(platform.AuthRequired, "credentials are unavailable", err)
}
var ret struct {
Code int `json:"code"`
Data struct {
IsLogin bool `json:"isLogin"`
} `json:"data"`
}
if err := a.getJSON(ctx, a.memberBase()+"/x/web-interface/nav", string(raw), &ret); err != nil {
return err
}
if ret.Code != 0 || !ret.Data.IsLogin {
return perr(platform.AuthRequired, "Bilibili login has expired", nil)
}
return nil
}
func (a *Adapter) Publish(ctx context.Context, in platform.PublishInput) (platform.PublishResult, error) {
s := a.store()
if s == nil {
return platform.PublishResult{}, perr(platform.Unknown, "credential store is not configured", nil)
}
raw, err := s.Load(in.WorkspaceID, in.AccountID)
if err != nil {
return platform.PublishResult{}, perr(platform.AuthRequired, "credentials are unavailable", err)
}
if len(in.MaterialURLs) == 0 {
return platform.PublishResult{}, perr(platform.Validation, "material is required", nil)
}
if len(in.MaterialURLs) > 1 {
return platform.PublishResult{}, perr(platform.Validation, "Bilibili currently accepts one video per task", nil)
}
path, cleanup, err := a.download(ctx, in.MaterialURLs[0])
if err != nil {
return platform.PublishResult{}, err
}
defer cleanup()
part, err := a.upload(ctx, string(raw), path)
if err != nil {
return platform.PublishResult{}, err
}
bvid, err := a.submit(ctx, string(raw), in, part)
if err != nil {
return platform.PublishResult{}, err
}
if bvid == "" {
return platform.PublishResult{}, perr(platform.PlatformChanged, "submission did not return bvid", nil)
}
return platform.PublishResult{URL: "https://www.bilibili.com/video/" + bvid, Receipt: bvid}, nil
}
func (a *Adapter) download(ctx context.Context, rawURL string) (string, func(), error) {
raw, err := a.request(ctx, http.MethodGet, rawURL, "", nil, nil)
if err != nil {
return "", nil, err
}
f, err := os.CreateTemp("", "ep-bilibili-*")
if err != nil {
return "", nil, perr(platform.Unknown, "creating temporary material", err)
}
name := f.Name()
if _, err = f.Write(raw); err != nil {
f.Close()
os.Remove(name)
return "", nil, perr(platform.Network, "saving material", err)
}
if err = f.Close(); err != nil {
os.Remove(name)
return "", nil, perr(platform.Network, "saving material", err)
}
return name, func() { _ = os.Remove(name) }, nil
}
func (a *Adapter) upload(ctx context.Context, cookies, path string) (map[string]interface{}, error) {
st, err := os.Stat(path)
if err != nil {
return nil, perr(platform.Validation, "material unavailable", err)
}
name, total := filepath.Base(path), st.Size()
if total <= 0 {
return nil, perr(platform.Validation, "material is empty", nil)
}
q := url.Values{"r": {"upos"}, "profile": {"ugcupos/bup"}, "ssl": {"0"}, "version": {"2.8.12"}, "build": {"2081200"}, "name": {name}, "size": {strconv.FormatInt(total, 10)}}
var pre struct {
OK int `json:"OK"`
Endpoint string `json:"endpoint"`
Auth string `json:"auth"`
BizID int `json:"biz_id"`
UposURI string `json:"upos_uri"`
ChunkSize int `json:"chunk_size"`
}
if err := a.getJSON(ctx, a.memberBase()+"/preupload?upcdn=bda2&probe_version=20221109&"+q.Encode(), cookies, &pre); err != nil {
return nil, err
}
if pre.OK != 1 || pre.Endpoint == "" || pre.UposURI == "" {
return nil, perr(platform.PlatformChanged, "invalid preupload response", nil)
}
upURL := a.upURL(pre.Endpoint, pre.UposURI)
var up struct {
UploadID string `json:"upload_id"`
}
if err := a.postJSON(ctx, upURL+"?uploads&output=json", cookies, map[string]string{"X-Upos-Auth": pre.Auth}, nil, &up); err != nil {
return nil, err
}
if up.UploadID == "" {
return nil, perr(platform.PlatformChanged, "upload id missing", nil)
}
chunk := pre.ChunkSize
if chunk <= 0 {
chunk = 7 * 1024 * 1024
}
count := int((total + int64(chunk) - 1) / int64(chunk))
f, err := os.Open(path)
if err != nil {
return nil, perr(platform.Validation, "material unavailable", err)
}
defer f.Close()
parts := make([]map[string]interface{}, 0, count)
buf := make([]byte, chunk)
for i := 0; i < count; i++ {
n, er := io.ReadFull(f, buf)
if er != nil && er != io.ErrUnexpectedEOF && er != io.EOF {
return nil, perr(platform.Network, "reading material", er)
}
pq := url.Values{"uploadId": {up.UploadID}, "partNumber": {strconv.Itoa(i + 1)}, "chunk": {strconv.Itoa(i)}, "chunks": {strconv.Itoa(count)}, "size": {strconv.Itoa(n)}, "start": {strconv.Itoa(i * chunk)}, "end": {strconv.Itoa(i*chunk + n)}, "total": {strconv.FormatInt(total, 10)}}
etag, err := a.put(ctx, upURL+"?"+pq.Encode(), cookies, map[string]string{"X-Upos-Auth": pre.Auth}, buf[:n])
if err != nil {
return nil, fmt.Errorf("%w: part %d", err, i+1)
}
if etag == "" {
return nil, perr(platform.PlatformChanged, "upload response missing ETag", nil)
}
parts = append(parts, map[string]interface{}{"partNumber": i + 1, "eTag": etag})
}
mq := url.Values{"name": {name}, "uploadId": {up.UploadID}, "biz_id": {strconv.Itoa(pre.BizID)}, "output": {"json"}, "profile": {"ugcupos/bup"}}
var merged struct {
OK int `json:"OK"`
}
if err := a.postJSON(ctx, upURL+"?"+mq.Encode(), cookies, map[string]string{"X-Upos-Auth": pre.Auth}, map[string]interface{}{"parts": parts}, &merged); err != nil {
return nil, err
}
if merged.OK != 1 {
return nil, perr(platform.PlatformChanged, "merge failed", nil)
}
base := strings.TrimSuffix(filepath.Base(pre.UposURI), filepath.Ext(filepath.Base(pre.UposURI)))
return map[string]interface{}{"title": strings.TrimSuffix(name, filepath.Ext(name)), "filename": base, "desc": ""}, nil
}
func (a *Adapter) submit(ctx context.Context, cookies string, in platform.PublishInput, part map[string]interface{}) (string, error) {
tag := strings.Join(in.Tags, ",")
if tag == "" {
tag = "日常"
}
tid := in.CategoryID
if tid <= 0 {
tid = 174
}
body := map[string]interface{}{"title": truncate(in.Title, 80), "desc": in.Content, "desc_v2": []map[string]interface{}{{"raw_text": in.Content, "type": 1, "biz_id": ""}}, "copyright": 1, "source": "", "tid": tid, "tag": tag, "dynamic": "", "videos": []map[string]interface{}{part}}
q := url.Values{"csrf": {cookieValue(cookies, "bili_jct")}}
var ret struct {
Code int `json:"code"`
Message string `json:"message"`
Data struct {
Bvid string `json:"bvid"`
} `json:"data"`
}
if err := a.postJSON(ctx, a.memberBase()+"/x/vu/web/add?"+q.Encode(), cookies, nil, body, &ret); err != nil {
return "", err
}
if ret.Code != 0 {
return "", perr(platform.Validation, ret.Message, nil)
}
return ret.Data.Bvid, nil
}
func (a *Adapter) getJSON(ctx context.Context, rawURL, cookies string, out interface{}) error {
raw, err := a.request(ctx, http.MethodGet, rawURL, cookies, nil, nil)
if err != nil {
return err
}
if err := json.Unmarshal(raw, out); err != nil {
return perr(platform.PlatformChanged, "invalid JSON response", err)
}
return nil
}
func (a *Adapter) postJSON(ctx context.Context, rawURL, cookies string, headers map[string]string, body, out interface{}) error {
var data io.Reader
if body != nil {
raw, err := json.Marshal(body)
if err != nil {
return perr(platform.Validation, "invalid request", err)
}
data = bytes.NewReader(raw)
}
raw, err := a.request(ctx, http.MethodPost, rawURL, cookies, headers, data)
if err != nil {
return err
}
if out != nil {
if err := json.Unmarshal(raw, out); err != nil {
return perr(platform.PlatformChanged, "invalid JSON response", err)
}
}
return nil
}
func (a *Adapter) put(ctx context.Context, rawURL, cookies string, headers map[string]string, body []byte) (string, error) {
for attempt := 0; attempt <= a.retries(); attempt++ {
if err := ctx.Err(); err != nil {
return "", err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPut, rawURL, bytes.NewReader(body))
if err != nil {
return "", perr(platform.Validation, "invalid URL", err)
}
if cookies != "" {
req.Header.Set("Cookie", cookies)
}
req.Header.Set("User-Agent", "Mozilla/5.0")
req.Header.Set("Referer", "https://member.bilibili.com")
for k, v := range headers {
req.Header.Set(k, v)
}
resp, err := a.client().Do(req)
if err != nil {
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return "", err
}
if attempt < a.retries() {
if !sleep(ctx, retryDelay(attempt)) {
return "", ctx.Err()
}
continue
}
return "", perr(platform.Network, "network request failed", err)
}
_, readErr := io.Copy(io.Discard, resp.Body)
resp.Body.Close()
if readErr != nil {
return "", perr(platform.Network, "reading response", readErr)
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
if resp.StatusCode >= 500 && attempt < a.retries() {
if !sleep(ctx, retryDelay(attempt)) {
return "", ctx.Err()
}
continue
}
cat := platform.Unknown
if resp.StatusCode == 401 || resp.StatusCode == 403 {
cat = platform.AuthRequired
}
if resp.StatusCode == 429 {
cat = platform.RateLimited
}
if resp.StatusCode >= 500 {
cat = platform.Network
}
return "", perr(cat, "HTTP "+strconv.Itoa(resp.StatusCode), nil)
}
return strings.Trim(resp.Header.Get("ETag"), "\""), nil
}
return "", perr(platform.Network, "upload failed", nil)
}
func (a *Adapter) request(ctx context.Context, method, rawURL, cookies string, headers map[string]string, body io.Reader) ([]byte, error) {
var payload []byte
var err error
if body != nil {
payload, err = io.ReadAll(body)
if err != nil {
return nil, perr(platform.Network, "reading request body", err)
}
}
var last error
for attempt := 0; attempt <= a.retries(); attempt++ {
if err := ctx.Err(); err != nil {
return nil, err
}
var reqBody io.Reader
if body != nil {
reqBody = bytes.NewReader(payload)
}
req, err := http.NewRequestWithContext(ctx, method, rawURL, reqBody)
if err != nil {
return nil, perr(platform.Validation, "invalid URL", err)
}
if cookies != "" {
req.Header.Set("Cookie", cookies)
}
req.Header.Set("User-Agent", "Mozilla/5.0")
req.Header.Set("Referer", "https://member.bilibili.com")
if body != nil && method == http.MethodPost {
req.Header.Set("Content-Type", "application/json")
}
for k, v := range headers {
req.Header.Set(k, v)
}
resp, err := a.client().Do(req)
if err != nil {
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return nil, err
}
last = perr(platform.Network, "network request failed", err)
if attempt < a.retries() {
if !sleep(ctx, retryDelay(attempt)) {
return nil, ctx.Err()
}
continue
}
return nil, last
}
raw, readErr := io.ReadAll(io.LimitReader(resp.Body, 16<<20))
resp.Body.Close()
if readErr != nil {
return nil, perr(platform.Network, "reading response", readErr)
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
cat := platform.Unknown
if resp.StatusCode == 401 || resp.StatusCode == 403 {
cat = platform.AuthRequired
}
if resp.StatusCode == 429 {
cat = platform.RateLimited
}
if resp.StatusCode >= 500 {
cat = platform.Network
}
last = perr(cat, "HTTP "+strconv.Itoa(resp.StatusCode), nil)
if resp.StatusCode >= 500 && attempt < a.retries() {
if !sleep(ctx, retryDelay(attempt)) {
return nil, ctx.Err()
}
continue
}
return nil, last
}
return raw, nil
}
return nil, last
}
func (a *Adapter) client() *http.Client {
if a.HTTP != nil {
return a.HTTP
}
return http.DefaultClient
}
func (a *Adapter) retries() int {
if a.MaxRetries < 0 {
return 0
}
if a.MaxRetries == 0 {
return 2
}
return a.MaxRetries
}
func retryDelay(n int) time.Duration { return time.Duration(25*(1<<n)) * time.Millisecond }
func sleep(ctx context.Context, d time.Duration) bool {
t := time.NewTimer(d)
defer t.Stop()
select {
case <-ctx.Done():
return false
case <-t.C:
return true
}
}
func (a *Adapter) memberBase() string {
if a.MemberBase != "" {
return strings.TrimRight(a.MemberBase, "/")
}
return "https://member.bilibili.com"
}
func (a *Adapter) passportBase() string {
if a.PassportBase != "" {
return strings.TrimRight(a.PassportBase, "/")
}
return "https://passport.bilibili.com"
}
func (a *Adapter) store() *credentials.Store {
if a.Store != nil {
return a.Store
}
return a.Credentials
}
func (a *Adapter) upURL(endpoint, uri string) string {
scheme := a.UpOSScheme
if scheme == "" {
scheme = "https"
}
if strings.HasPrefix(endpoint, "http://") || strings.HasPrefix(endpoint, "https://") {
return strings.TrimRight(endpoint, "/") + "/" + strings.TrimPrefix(uri, "upos://")
}
return scheme + "://" + strings.TrimRight(strings.TrimPrefix(endpoint, "//"), "/") + "/" + strings.TrimPrefix(uri, "upos://")
}
// cookiesFromURL extracts only the three cookies accepted by Bilibili's
// cross-domain login redirect; all query values are URL-decoded by net/url.
func cookiesFromURL(raw string) string {
u, err := url.Parse(raw)
if err != nil {
return ""
}
q := u.Query()
sess := q.Get("SESSDATA")
csrf := q.Get("bili_jct")
uid := q.Get("DedeUserID")
if sess == "" || csrf == "" || uid == "" {
return ""
}
return "SESSDATA=" + sess + "; bili_jct=" + csrf + "; DedeUserID=" + uid
}
func cookieValue(cookies, key string) string {
for _, p := range strings.Split(cookies, ";") {
p = strings.TrimSpace(p)
i := strings.IndexByte(p, '=')
if i > 0 && p[:i] == key {
return p[i+1:]
}
}
return ""
}
func truncate(s string, n int) string {
r := []rune(s)
if len(r) > n {
return string(r[:n])
}
return s
}
func perr(cat platform.ErrorCategory, msg string, err error) error {
return &platform.Error{Category: cat, Message: msg, Err: err}
}
@@ -0,0 +1,145 @@
package bilibili
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"everypublish/server/internal/credentials"
"everypublish/server/internal/platform"
)
func TestQRLoginAndCredentialSave(t *testing.T) {
var poll int
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
if strings.HasSuffix(r.URL.Path, "/generate") {
_, _ = w.Write([]byte(`{"code":0,"data":{"url":"https://qr.test/x","qrcode_key":"k"}}`))
return
}
poll++
switch poll {
case 1:
_, _ = w.Write([]byte(`{"code":86101}`))
case 2:
_, _ = w.Write([]byte(`{"code":86090}`))
default:
_, _ = w.Write([]byte(`{"code":0,"data":{"url":"https://passport.test/c?SESSDATA=s%2Bv&bili_jct=j%2Fv&DedeUserID=7"}}`))
}
}))
defer srv.Close()
dir := t.TempDir()
store, err := credentials.Open(dir)
if err != nil {
t.Fatal(err)
}
a := New(store)
a.PassportBase = srv.URL
s, err := a.BeginLogin(context.Background())
if err != nil || s.Token != "k" {
t.Fatalf("begin: %#v %v", s, err)
}
p, err := a.PollLogin(context.Background(), s.Token)
if err != nil || p.State != platform.LoginPending {
t.Fatalf("pending: %#v %v", p, err)
}
p, err = a.PollLogin(context.Background(), s.Token)
if err != nil || p.State != platform.LoginScanned {
t.Fatalf("scanned: %#v %v", p, err)
}
p, err = a.PollLogin(context.Background(), s.Token)
if err != nil || p.Cookies != "SESSDATA=s+v; bili_jct=j/v; DedeUserID=7" {
t.Fatalf("confirmed: %#v %v", p, err)
}
if err := a.SaveCredentials(1, 2, p.Cookies); err != nil {
t.Fatal(err)
}
got, err := store.Load(1, 2)
if err != nil || string(got) != p.Cookies {
t.Fatalf("stored=%q err=%v", got, err)
}
}
func TestPublishUPOSUsesETagAndEncodedCSRF(t *testing.T) {
dir := t.TempDir()
store, err := credentials.Open(dir)
if err != nil {
t.Fatal(err)
}
if err = store.Save(1, 2, []byte("SESSDATA=s; bili_jct=a&b; DedeUserID=7")); err != nil {
t.Fatal(err)
}
var sawPart, sawMerge, sawSubmit bool
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch r.URL.Path {
case "/material":
_, _ = w.Write([]byte("video"))
case "/preupload":
_, _ = w.Write([]byte(`{"OK":1,"endpoint":"//` + r.Host + `","auth":"auth","biz_id":1,"upos_uri":"upos://folder/v.mp4","chunk_size":2}`))
case "/folder/v.mp4":
if r.Method == http.MethodPost && r.URL.Query().Has("uploads") {
_, _ = w.Write([]byte(`{"upload_id":"u1"}`))
} else if r.Method == http.MethodPut {
sawPart = true
w.Header().Set("ETag", `"real-etag"`)
w.WriteHeader(http.StatusOK)
} else if r.Method == http.MethodPost {
sawMerge = true
_, _ = w.Write([]byte(`{"OK":1}`))
}
case "/x/vu/web/add":
if r.URL.Query().Get("csrf") != "a&b" {
t.Errorf("csrf not decoded: %q", r.URL.RawQuery)
}
sawSubmit = true
_, _ = w.Write([]byte(`{"code":0,"data":{"bvid":"BV1"}}`))
}
}))
defer srv.Close()
a := New(store)
a.MemberBase = srv.URL
a.UpOSScheme = "http"
material := srv.URL + "/material"
result, err := a.Publish(context.Background(), platform.PublishInput{WorkspaceID: 1, AccountID: 2, Title: "title", Content: "body", MaterialURLs: []string{material}})
if err != nil {
t.Fatal(err)
}
if result.URL != "https://www.bilibili.com/video/BV1" || !sawPart || !sawMerge || !sawSubmit {
t.Fatalf("result=%#v part=%v merge=%v submit=%v", result, sawPart, sawMerge, sawSubmit)
}
}
func TestSaveCredentialsRejectsIncompleteCookie(t *testing.T) {
store, err := credentials.Open(t.TempDir())
if err != nil {
t.Fatal(err)
}
if err := New(store).SaveCredentials(1, 1, "SESSDATA=only"); err == nil || platform.CategoryOf(err) != platform.Validation {
t.Fatalf("expected validation error for incomplete cookie, got %v", err)
}
}
func TestCheckCredentialsUsesEncryptedCookieAndNavEndpoint(t *testing.T) {
store, err := credentials.Open(t.TempDir())
if err != nil {
t.Fatal(err)
}
if err = store.Save(1, 2, []byte("SESSDATA=s; bili_jct=j; DedeUserID=7")); err != nil {
t.Fatal(err)
}
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/x/web-interface/nav" || r.Header.Get("Cookie") != "SESSDATA=s; bili_jct=j; DedeUserID=7" {
t.Fatalf("unexpected nav request: %s cookie=%q", r.URL.Path, r.Header.Get("Cookie"))
}
_, _ = w.Write([]byte(`{"code":0,"data":{"isLogin":true}}`))
}))
defer srv.Close()
a := New(store)
a.MemberBase = srv.URL
if err := a.CheckCredentials(context.Background(), 1, 2); err != nil {
t.Fatal(err)
}
}
+146
View File
@@ -0,0 +1,146 @@
package platform
import (
"context"
"errors"
"fmt"
"sort"
"sync"
"time"
)
// LoginState is deliberately small: the web UI maps it to pending/scanned/
// solved without knowing platform-specific numeric QR codes.
type LoginState string
const (
LoginPending LoginState = "pending"
LoginScanned LoginState = "scanned"
LoginConfirmed LoginState = "confirmed"
)
type LoginSession struct {
Token string
URL string
}
type LoginPoll struct {
State LoginState
Cookies string
Prompt string
}
type PublishInput struct {
WorkspaceID uint64
AccountID uint64
TaskID uint64
Title string
Content string
Tags []string
MaterialURLs []string
CategoryID int
ScheduleAt *time.Time
}
type PublishResult struct {
URL string
Receipt string
}
type ErrorCategory string
const (
AuthRequired ErrorCategory = "auth_required"
ChallengeRequired ErrorCategory = "challenge_required"
Validation ErrorCategory = "validation"
RateLimited ErrorCategory = "rate_limited"
Network ErrorCategory = "network"
PlatformChanged ErrorCategory = "platform_changed"
Unknown ErrorCategory = "unknown"
)
// Error is safe to show in task history: it carries a category and excludes
// cookies, tokens and raw response bodies by construction.
type Error struct {
Category ErrorCategory
Message string
Err error
}
func (e *Error) Error() string {
if e == nil {
return ""
}
if e.Message != "" {
return fmt.Sprintf("%s: %s", e.Category, e.Message)
}
return string(e.Category)
}
func (e *Error) Unwrap() error { return e.Err }
func CategoryOf(err error) ErrorCategory {
var pe *Error
if errors.As(err, &pe) && pe.Category != "" {
return pe.Category
}
return Unknown
}
// Adapter is the server-side boundary for a platform. Implementations own
// cookies/profile details; API handlers only persist challenge state.
type Adapter interface {
Platform() string
BeginLogin(context.Context) (LoginSession, error)
PollLogin(context.Context, string) (LoginPoll, error)
SaveCredentials(workspaceID, accountID uint64, cookies string) error
Publish(context.Context, PublishInput) (PublishResult, error)
}
// CredentialChecker is optional so a platform can expose a cheap login health
// probe without forcing every adapter to implement a platform-specific API.
type CredentialChecker interface {
CheckCredentials(context.Context, uint64, uint64) error
}
type Registry struct {
mu sync.RWMutex
adapters map[string]Adapter
}
func NewRegistry() *Registry { return &Registry{adapters: make(map[string]Adapter)} }
func (r *Registry) Register(adapter Adapter) error {
if r == nil || adapter == nil || adapter.Platform() == "" {
return errors.New("invalid platform adapter")
}
r.mu.Lock()
defer r.mu.Unlock()
r.adapters[adapter.Platform()] = adapter
return nil
}
func (r *Registry) Get(name string) Adapter {
if r == nil {
return nil
}
r.mu.RLock()
defer r.mu.RUnlock()
return r.adapters[name]
}
func (r *Registry) Has(name string) bool { return r.Get(name) != nil }
func (r *Registry) Names() []string {
if r == nil {
return nil
}
r.mu.RLock()
defer r.mu.RUnlock()
names := make([]string, 0, len(r.adapters))
for name := range r.adapters {
names = append(names, name)
}
sort.Strings(names)
return names
}
+370 -86
View File
@@ -1,34 +1,50 @@
package ws
import (
"context"
"encoding/json"
"errors"
"fmt"
"log"
"strconv"
"net/url"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/config"
"everypublish/server/internal/credentials"
"everypublish/server/internal/models"
"everypublish/server/internal/platform"
"everypublish/server/internal/task"
"everypublish/shared/proto"
)
// Dispatcher 任务下发器(MVP:进程内轮询 + 审批/重试即时触发;接口预留二期换 asynq)
// Dispatcher 任务执行器(Web-only 本机 worker;审批/重试即时触发)
type Dispatcher struct {
DB *gorm.DB
Hub *Hub
Cfg *config.Config
Interval time.Duration
DB *gorm.DB
Hub *Hub
Cfg *config.Config
Registry *platform.Registry
Credentials *credentials.Store
Interval time.Duration
runsMu sync.Mutex
runs map[uint64]context.CancelFunc
}
type publishTarget struct {
account models.Account
adapter platform.Adapter
}
// NewDispatcher 构造下发器
func NewDispatcher(db *gorm.DB, hub *Hub, cfg *config.Config) *Dispatcher {
return &Dispatcher{DB: db, Hub: hub, Cfg: cfg, Interval: time.Second}
return &Dispatcher{DB: db, Hub: hub, Cfg: cfg, Interval: time.Second, runs: make(map[uint64]context.CancelFunc)}
}
// Start 后台轮询:把到期且排队的任务下发给在线设备
// Start 后台轮询:把到期且排队的任务交给当前执行模式
func (d *Dispatcher) Start(stop <-chan struct{}) {
if d.Interval <= 0 {
d.Interval = time.Second
@@ -57,7 +73,7 @@ func (d *Dispatcher) poll() {
}
}
// TryDispatch 立即下发单个任务(审批/重试后调用);无在线设备返回 false
// TryDispatch 立即执行单个任务(审批/重试后调用)。
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 {
@@ -66,96 +82,364 @@ func (d *Dispatcher) TryDispatch(taskID uint64) bool {
if row.ScheduleAt != nil && row.ScheduleAt.After(time.Now()) {
return false
}
devID := d.pickDevice(&row)
if devID == 0 {
if d.Cfg != nil && d.Cfg.ExecutorMode != "mock" && d.Registry != nil {
ids := decodeIDs(row.AccountIDs)
targets := make([]publishTarget, 0, len(ids))
allRegistered := len(ids) > 0
hasRegistered := false
for _, id := range ids {
var account models.Account
if d.DB.Where("id = ? and workspace_id = ?", id, row.WorkspaceID).First(&account).Error != nil {
allRegistered = false
continue
}
adapter := d.Registry.Get(account.Platform)
if adapter == nil {
allRegistered = false
continue
}
hasRegistered = true
targets = append(targets, publishTarget{account: account, adapter: adapter})
}
if allRegistered {
return d.dispatchPlatform(&row, targets)
}
if hasRegistered {
return d.failUnsupportedPlatform(&row)
}
}
return d.dispatchLocalMock(&row)
}
func (d *Dispatcher) failUnsupportedPlatform(row *models.Task) bool {
next, ok := task.Next(row.Status, task.ActionDispatch)
if !ok || d.DB.Model(row).Update("status", next).Error != nil {
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
d.recordEvent(row, string(task.ActionDispatch), "任务包含尚未接入真实适配器的平台")
d.broadcastTask(*row)
if next, ok = task.Next(row.Status, task.ActionFail); ok {
message := "任务包含尚未接入真实适配器的平台,请拆分任务后重试"
_ = d.DB.Model(row).Updates(map[string]interface{}{"status": next, "error_message": message}).Error
row.Status = next
row.ErrorMessage = message
d.recordEvent(row, string(task.ActionFail), message)
d.broadcastTask(*row)
}
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 隔离
}
// dispatchPlatform runs a registered adapter on the server itself. A task is
// marked dispatched before the goroutine starts so retries/cancel operations
// observe the same state machine as the local mock executor.
func (d *Dispatcher) dispatchPlatform(row *models.Task, targets []publishTarget) bool {
next, ok := task.Next(row.Status, task.ActionDispatch)
if !ok {
return false
}
online := d.Hub.OnlineDevices(row.WorkspaceID)
if len(online) > 0 {
return online[0]
claim := d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Queued).Update("status", next)
if claim.Error != nil || claim.RowsAffected != 1 {
return false
}
return 0
row.Status = next
d.recordEvent(row, string(task.ActionDispatch), "平台适配器已接收任务")
d.broadcastTask(*row)
ctx, cancel := context.WithCancel(context.Background())
d.runsMu.Lock()
if d.runs == nil {
d.runs = make(map[uint64]context.CancelFunc)
}
d.runs[row.ID] = cancel
d.runsMu.Unlock()
go d.runPlatform(ctx, row.ID, row.WorkspaceID, targets)
return true
}
// 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)
func (d *Dispatcher) runPlatform(ctx context.Context, taskID, workspaceID uint64, targets []publishTarget) {
defer func() {
d.runsMu.Lock()
delete(d.runs, taskID)
d.runsMu.Unlock()
}()
var row models.Task
if err := d.DB.Where("id = ? and workspace_id = ?", taskID, workspaceID).First(&row).Error; err != nil || row.Status != task.Dispatched {
return
}
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 next, ok := task.Next(row.Status, task.ActionStart); ok {
claim := d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Dispatched).Update("status", next)
if claim.Error != nil || claim.RowsAffected != 1 {
return
}
row.Status = next
d.recordEvent(&row, string(task.ActionStart), "平台适配器开始处理")
d.broadcastTask(row)
}
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
var err error
results := make([]platform.PublishResult, 0, len(targets))
for _, target := range targets {
input, inputErr := d.publishInput(row, target.account)
if inputErr != nil {
err = inputErr
break
}
result, publishErr := target.adapter.Publish(ctx, input)
if publishErr != nil {
err = publishErr
break
}
if strings.TrimSpace(result.URL) == "" {
err = &platform.Error{Category: platform.PlatformChanged, Message: "平台未返回发布链接"}
break
}
results = append(results, result)
}
if err == nil && len(results) == len(targets) {
urlsList := make([]string, 0, len(results))
receiptsList := make([]string, 0, len(results))
for _, result := range results {
urlsList = append(urlsList, result.URL)
receiptsList = append(receiptsList, result.Receipt)
}
urls, _ := json.Marshal(urlsList)
receipts, _ := json.Marshal(receiptsList)
now := time.Now()
tx := d.DB.Begin()
committed := false
if tx.Error == nil {
result := tx.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Running).Updates(map[string]interface{}{
"status": task.Success, "published_urls": string(urls), "receipts": string(receipts),
"error_message": "", "published_at": now,
})
txErr := result.Error
if txErr == nil && result.RowsAffected != 1 {
txErr = context.Canceled
}
if txErr == nil {
txErr = tx.Create(&models.TaskEvent{WorkspaceID: row.WorkspaceID, TaskID: row.ID, Type: string(task.ActionSuccess), Status: task.Success, Message: "平台发布完成"}).Error
}
if txErr == nil {
txErr = tx.Commit().Error
}
if txErr != nil {
tx.Rollback()
} else {
committed = true
}
}
if committed {
row.Status = task.Success
row.PublishedURLs = string(urls)
row.Receipts = string(receipts)
row.PublishedAt = &now
d.notifySuccess(row)
d.broadcastTask(row)
return
}
}
if row.ScheduleAt != nil {
push.ScheduleAt = row.ScheduleAt.UnixMilli()
message := "平台发布失败"
if err != nil {
message = safePlatformError(err)
}
var current models.Task
if d.DB.Where("id = ? and workspace_id = ?", row.ID, row.WorkspaceID).First(&current).Error == nil && current.Status == task.Cancelled {
return
}
if next, ok := task.Next(row.Status, task.ActionFail); ok {
updates := map[string]interface{}{"status": next, "error_message": message}
if len(results) > 0 {
partialURLs := make([]string, 0, len(results))
partialReceipts := make([]string, 0, len(results))
for _, result := range results {
partialURLs = append(partialURLs, result.URL)
partialReceipts = append(partialReceipts, result.Receipt)
}
partialURLJSON, _ := json.Marshal(partialURLs)
partialReceiptJSON, _ := json.Marshal(partialReceipts)
updates["published_urls"] = string(partialURLJSON)
updates["receipts"] = string(partialReceiptJSON)
message = "部分账号已发布;" + message
updates["error_message"] = message
}
_ = d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Running).Updates(updates).Error
row.Status = next
row.ErrorMessage = message
d.recordEvent(&row, string(task.ActionFail), message)
d.broadcastTask(row)
}
}
// Cancel interrupts a real adapter run. The task handler still performs the
// persisted state transition so the database and worker agree on cancellation.
func (d *Dispatcher) Cancel(taskID uint64) {
d.runsMu.Lock()
cancel := d.runs[taskID]
d.runsMu.Unlock()
if cancel != nil {
cancel()
}
}
func (d *Dispatcher) publishInput(row models.Task, account models.Account) (platform.PublishInput, error) {
materialIDs := decodeIDs(row.MaterialIDs)
if len(materialIDs) == 0 {
return platform.PublishInput{}, &platform.Error{Category: platform.Validation, Message: "任务没有素材"}
}
var materials []models.Material
if err := d.DB.Where("workspace_id = ? and id in ? and status = ?", row.WorkspaceID, materialIDs, "ready").Find(&materials).Error; err != nil {
return platform.PublishInput{}, &platform.Error{Category: platform.Unknown, Message: "读取素材失败", Err: err}
}
if len(materials) != len(materialIDs) {
return platform.PublishInput{}, &platform.Error{Category: platform.Validation, Message: "任务素材不存在或未就绪"}
}
urls := make([]string, 0, len(materials))
for _, material := range materials {
token := models.TransferToken{MaterialID: material.ID, Token: uuid.NewString(), ExpiresAt: time.Now().Add(10 * time.Minute)}
if err := d.DB.Create(&token).Error; err != nil {
return platform.PublishInput{}, &platform.Error{Category: platform.Unknown, Message: "生成素材访问链接失败", Err: err}
}
base := ""
if d.Cfg != nil {
base = strings.TrimRight(d.Cfg.BaseURL, "/")
}
if base == "" {
base = "http://127.0.0.1:8090"
}
urls = append(urls, base+"/api/v1/files/"+url.QueryEscape(token.Token))
}
var tags []string
if row.Tags != "" {
_ = json.Unmarshal([]byte(row.Tags), &tags)
}
return platform.PublishInput{WorkspaceID: row.WorkspaceID, AccountID: account.ID, TaskID: row.ID, Title: row.Title, Content: row.Content, Tags: tags, MaterialURLs: urls, ScheduleAt: row.ScheduleAt}, nil
}
func (d *Dispatcher) notifySuccess(row models.Task) {
if row.CreatedBy != 0 {
_ = d.DB.Create(&models.Notification{WorkspaceID: row.WorkspaceID, UserID: row.CreatedBy, Kind: "task", Title: "任务执行完成", Content: row.Title}).Error
}
if d.Hub != nil {
d.Hub.BroadcastWS(row.WorkspaceID, "notification.new", gin.H{"kind": "task", "title": "任务执行完成", "content": row.Title})
}
}
func decodeIDs(raw string) []uint64 {
var ids []uint64
if err := json.Unmarshal([]byte(raw), &ids); err != nil {
return nil
}
return ids
}
func safePlatformError(err error) string {
var pe *platform.Error
if errors.As(err, &pe) {
return pe.Error()
}
return "平台发布失败"
}
// dispatchLocalMock is the Web-only V1 executor. It deliberately produces a
// mock:// receipt so a local test can prove the state machine without claiming
// that a real platform accepted the content.
func (d *Dispatcher) dispatchLocalMock(row *models.Task) bool {
next, ok := task.Next(row.Status, task.ActionDispatch)
if !ok {
return false
}
claim := d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", row.ID, row.WorkspaceID, task.Queued).Update("status", next)
if claim.Error != nil || claim.RowsAffected != 1 {
return false
}
row.Status = next
d.recordEvent(row, string(task.ActionDispatch), "本机 Web 执行器已接收任务")
d.broadcastTask(*row)
go func(taskID uint64, workspaceID uint64) {
time.Sleep(120 * time.Millisecond)
var running models.Task
if err := d.DB.Where("id = ? and workspace_id = ?", taskID, workspaceID).First(&running).Error; err != nil || running.Status != task.Dispatched {
return
}
if next, ok := task.Next(running.Status, task.ActionStart); ok {
claim := d.DB.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", running.ID, running.WorkspaceID, task.Dispatched).Update("status", next)
if claim.Error != nil || claim.RowsAffected != 1 {
return
}
running.Status = next
d.recordEvent(&running, string(task.ActionStart), "本机 Web 执行器开始处理")
d.broadcastTask(running)
}
time.Sleep(180 * time.Millisecond)
if err := d.DB.Where("id = ? and workspace_id = ?", taskID, workspaceID).First(&running).Error; err != nil || running.Status != task.Running {
return
}
next, ok := task.Next(running.Status, task.ActionSuccess)
if !ok {
return
}
urls, _ := json.Marshal([]string{fmt.Sprintf("mock://everypublish/task/%d", taskID)})
receipts, _ := json.Marshal([]string{"mock executor completed"})
now := time.Now()
tx := d.DB.Begin()
if tx.Error != nil {
return
}
result := tx.Model(&models.Task{}).Where("id = ? and workspace_id = ? and status = ?", running.ID, running.WorkspaceID, task.Running).Updates(map[string]interface{}{
"status": next, "published_urls": string(urls), "receipts": string(receipts),
"error_message": "", "published_at": now,
})
if result.Error != nil || result.RowsAffected != 1 {
tx.Rollback()
return
}
running.Status = next
running.PublishedURLs = string(urls)
running.Receipts = string(receipts)
running.PublishedAt = &now
if err := tx.Create(&models.TaskEvent{
WorkspaceID: running.WorkspaceID,
TaskID: running.ID,
Type: string(task.ActionSuccess),
Status: running.Status,
Message: "本机 mock 执行完成",
}).Error; err != nil {
tx.Rollback()
return
}
if err := tx.Commit().Error; err != nil {
return
}
if running.CreatedBy != 0 {
_ = d.DB.Create(&models.Notification{
WorkspaceID: running.WorkspaceID,
UserID: running.CreatedBy,
Kind: "task",
Title: "任务执行完成",
Content: running.Title,
}).Error
}
if d.Hub != nil {
d.Hub.BroadcastWS(running.WorkspaceID, "notification.new", gin.H{"kind": "task", "title": "任务执行完成", "content": running.Title})
}
d.broadcastTask(running)
}(row.ID, row.WorkspaceID)
return true
}
func (d *Dispatcher) recordEvent(row *models.Task, eventType, message string) {
_ = d.DB.Create(&models.TaskEvent{
WorkspaceID: row.WorkspaceID,
TaskID: row.ID,
Type: eventType,
Status: row.Status,
Message: message,
}).Error
}
func (d *Dispatcher) broadcastTask(row models.Task) {
if d.Hub != nil {
d.Hub.BroadcastWS(row.WorkspaceID, "task.status", taskEvent(row))
}
return push
}
@@ -0,0 +1,46 @@
package ws
import (
"testing"
"time"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
"everypublish/server/internal/task"
gormsqlite "gorm.io/driver/sqlite"
"gorm.io/gorm"
)
func TestDispatcherLocalMock(t *testing.T) {
db, err := gorm.Open(gormsqlite.Open("file:dispatcher_mock?mode=memory&cache=shared"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err = db.AutoMigrate(&models.Task{}, &models.TaskEvent{}); err != nil {
t.Fatal(err)
}
row := models.Task{WorkspaceID: 1, Title: "web-only", Status: task.Queued, AccountIDs: "[]", MaterialIDs: "[]"}
if err = db.Create(&row).Error; err != nil {
t.Fatal(err)
}
d := NewDispatcher(db, NewHub(), &config.Config{ExecutorMode: "mock"})
if !d.TryDispatch(row.ID) {
t.Fatal("mock dispatch should accept queued task")
}
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
var got models.Task
if err = db.First(&got, row.ID).Error; err != nil {
t.Fatal(err)
}
if got.Status == task.Success {
if got.PublishedURLs == "" || got.Receipts == "" {
t.Fatal("mock success must include URL and receipt")
}
return
}
time.Sleep(20 * time.Millisecond)
}
t.Fatal("mock task did not reach success")
}
@@ -0,0 +1,131 @@
package ws
import (
"context"
"path/filepath"
"sync"
"testing"
"time"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
"everypublish/server/internal/platform"
"everypublish/server/internal/task"
gormsqlite "gorm.io/driver/sqlite"
"gorm.io/gorm"
)
type testPlatformAdapter struct {
mu sync.Mutex
called int
block bool
canceled chan struct{}
}
func (a *testPlatformAdapter) Platform() string { return "test-platform" }
func (a *testPlatformAdapter) BeginLogin(context.Context) (platform.LoginSession, error) {
return platform.LoginSession{}, nil
}
func (a *testPlatformAdapter) PollLogin(context.Context, string) (platform.LoginPoll, error) {
return platform.LoginPoll{}, nil
}
func (a *testPlatformAdapter) SaveCredentials(uint64, uint64, string) error { return nil }
func (a *testPlatformAdapter) Publish(ctx context.Context, in platform.PublishInput) (platform.PublishResult, error) {
a.mu.Lock()
a.called++
a.mu.Unlock()
if a.block {
select {
case <-ctx.Done():
close(a.canceled)
return platform.PublishResult{}, ctx.Err()
case <-time.After(2 * time.Second):
}
}
return platform.PublishResult{URL: "https://example.test/task/" + in.Title, Receipt: "receipt"}, nil
}
func platformTestDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(gormsqlite.Open(filepath.Join(t.TempDir(), "dispatcher.db")), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err = db.AutoMigrate(&models.Task{}, &models.TaskEvent{}, &models.Account{}, &models.Material{}, &models.TransferToken{}, &models.Notification{}); err != nil {
t.Fatal(err)
}
return db
}
func TestDispatcherUsesRegisteredPlatformAndPublishesPerAccount(t *testing.T) {
db := platformTestDB(t)
account := models.Account{WorkspaceID: 1, Platform: "test-platform", Status: "active"}
material := models.Material{WorkspaceID: 1, Name: "video.mp4", Status: "ready", StorageKey: "1/video.mp4"}
if err := db.Create(&account).Error; err != nil {
t.Fatal(err)
}
if err := db.Create(&material).Error; err != nil {
t.Fatal(err)
}
row := models.Task{WorkspaceID: 1, Title: "platform", Status: task.Queued, AccountIDs: "[1]", MaterialIDs: "[1]"}
if err := db.Create(&row).Error; err != nil {
t.Fatal(err)
}
adapter := &testPlatformAdapter{}
registry := platform.NewRegistry()
if err := registry.Register(adapter); err != nil {
t.Fatal(err)
}
d := NewDispatcher(db, NewHub(), &config.Config{ExecutorMode: "web", BaseURL: "http://127.0.0.1:8090"})
d.Registry = registry
if !d.TryDispatch(row.ID) {
t.Fatal("registered adapter should accept queued task")
}
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
var got models.Task
if err := db.First(&got, row.ID).Error; err != nil {
t.Fatal(err)
}
if got.Status == task.Success {
if got.PublishedURLs != `["https://example.test/task/platform"]` || got.Receipts != `["receipt"]` {
t.Fatalf("unexpected platform result: urls=%s receipts=%s", got.PublishedURLs, got.Receipts)
}
adapter.mu.Lock()
called := adapter.called
adapter.mu.Unlock()
if called != 1 {
t.Fatalf("adapter called %d times", called)
}
return
}
time.Sleep(20 * time.Millisecond)
}
t.Fatal("platform task did not reach success")
}
func TestDispatcherCancelStopsRegisteredPlatform(t *testing.T) {
db := platformTestDB(t)
account := models.Account{WorkspaceID: 1, Platform: "test-platform", Status: "active"}
material := models.Material{WorkspaceID: 1, Name: "video.mp4", Status: "ready", StorageKey: "1/video.mp4"}
_ = db.Create(&account).Error
_ = db.Create(&material).Error
row := models.Task{WorkspaceID: 1, Title: "cancel", Status: task.Queued, AccountIDs: "[1]", MaterialIDs: "[1]"}
_ = db.Create(&row).Error
adapter := &testPlatformAdapter{block: true, canceled: make(chan struct{})}
registry := platform.NewRegistry()
_ = registry.Register(adapter)
d := NewDispatcher(db, NewHub(), &config.Config{ExecutorMode: "web", BaseURL: "http://127.0.0.1:8090"})
d.Registry = registry
if !d.TryDispatch(row.ID) {
t.Fatal("registered adapter should accept queued task")
}
time.Sleep(60 * time.Millisecond)
d.Cancel(row.ID)
select {
case <-adapter.canceled:
case <-time.After(time.Second):
t.Fatal("adapter context was not cancelled")
}
}