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
+5
View File
@@ -5,4 +5,9 @@ REDIS_PASSWORD=
JWT_SECRET=please-change-me-in-production
BASE_URL=http://127.0.0.1:8090
STORAGE_DIR=./data/materials
CREDENTIAL_DIR=./data/credentials
EXECUTOR_MODE=mock
BILIBILI_MEMBER_BASE=https://member.bilibili.com
BILIBILI_PASSPORT_BASE=https://passport.bilibili.com
BILIBILI_UPOS_SCHEME=https
STATIC_DIR=../apps/web/dist
+16
View File
@@ -48,6 +48,7 @@ CREATE TABLE IF NOT EXISTS accounts (
workspace_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
platform VARCHAR(32) NOT NULL DEFAULT '',
account_name VARCHAR(128) NOT NULL DEFAULT '',
remark VARCHAR(255) NOT NULL DEFAULT '',
avatar_url VARCHAR(512) NOT NULL DEFAULT '',
status VARCHAR(16) NOT NULL DEFAULT 'unbound',
health INT NOT NULL DEFAULT 100,
@@ -109,6 +110,21 @@ CREATE TABLE IF NOT EXISTS tasks (
KEY idx_tasks_schedule_at (schedule_at)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS task_events (
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT,
created_at DATETIME(3) NULL,
updated_at DATETIME(3) NULL,
workspace_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
task_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
type VARCHAR(32) NOT NULL DEFAULT '',
status VARCHAR(32) NOT NULL DEFAULT '',
message VARCHAR(1024) NOT NULL DEFAULT '',
data TEXT NULL,
PRIMARY KEY (id),
KEY idx_task_events_workspace_id (workspace_id),
KEY idx_task_events_task_id (task_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS agent_devices (
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT,
created_at DATETIME(3) NULL,
+2
View File
@@ -12,6 +12,7 @@ require (
github.com/redis/go-redis/v9 v9.22.0
golang.org/x/crypto v0.55.0
gorm.io/driver/mysql v1.6.0
gorm.io/driver/sqlite v1.6.0
gorm.io/gorm v1.31.2
)
@@ -36,6 +37,7 @@ require (
github.com/klauspost/cpuid/v2 v2.3.0 // indirect
github.com/leodido/go-urn v1.4.0 // indirect
github.com/mattn/go-isatty v0.0.20 // indirect
github.com/mattn/go-sqlite3 v1.14.22 // indirect
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
github.com/modern-go/reflect2 v1.0.2 // indirect
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
+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")
}
}
+18 -5
View File
@@ -15,8 +15,11 @@ import (
"everypublish/server/internal/api"
"everypublish/server/internal/cache"
"everypublish/server/internal/config"
"everypublish/server/internal/credentials"
"everypublish/server/internal/db"
"everypublish/server/internal/models"
"everypublish/server/internal/platform"
"everypublish/server/internal/platform/bilibili"
"everypublish/server/internal/ws"
)
@@ -41,18 +44,28 @@ func main() {
}
log.Println("redis connected")
serverKey, err := ws.LoadOrCreateServerKey("./data/server_ed25519.key")
if err != nil {
log.Fatalf("server ed25519 key failed: %v", err)
}
hub := ws.NewHub()
credentialStore, err := credentials.Open(cfg.CredentialDir)
if err != nil {
log.Fatalf("credential store failed: %v", err)
}
registry := platform.NewRegistry()
biliAdapter := bilibili.New(credentialStore)
biliAdapter.MemberBase = cfg.BilibiliMemberBase
biliAdapter.PassportBase = cfg.BilibiliPassportBase
biliAdapter.UpOSScheme = cfg.BilibiliUpOSScheme
if err := registry.Register(biliAdapter); err != nil {
log.Fatalf("register bilibili adapter failed: %v", err)
}
dsp := ws.NewDispatcher(gdb, hub, cfg)
dsp.Registry = registry
dsp.Credentials = credentialStore
ctx, stopFn := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stopFn()
go dsp.Start(ctx.Done())
router := api.Router(cfg, gdb, rds, hub, serverKey, dsp)
router := api.Router(cfg, gdb, rds, hub, dsp)
// 托管前端生产构建(STATIC_DIR 存在 index.html 时启用),SPA 回退
if cfg.StaticDir != "" {