server: 平台抽象层+bilibili、服务端加密凭据库、角色中间件;web: 精简页面/路由、macOS 客户端(Swift)与多份方案文档;移除误入库的编译产物
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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})
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)),
|
||||
|
||||
@@ -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"),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -43,6 +43,7 @@ func AutoMigrate(db *gorm.DB) error {
|
||||
&models.Account{},
|
||||
&models.Material{},
|
||||
&models.Task{},
|
||||
&models.TaskEvent{},
|
||||
&models.AgentDevice{},
|
||||
&models.Challenge{},
|
||||
&models.AuditLog{},
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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(¤t).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
@@ -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 != "" {
|
||||
|
||||
Reference in New Issue
Block a user