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

This commit is contained in:
Qiufeng
2026-08-20 20:38:21 +08:00
parent bf25d8beac
commit cc8845dcec
44 changed files with 4995 additions and 0 deletions
+8
View File
@@ -0,0 +1,8 @@
SERVER_ADDR=:8090
MYSQL_DSN=root:everypublish@tcp(127.0.0.1:3306)/everypublish?charset=utf8mb4&parseTime=True&loc=Local
REDIS_ADDR=127.0.0.1:6379
REDIS_PASSWORD=
JWT_SECRET=please-change-me-in-production
BASE_URL=http://127.0.0.1:8090
STORAGE_DIR=./data/materials
STATIC_DIR=../apps/web/dist
+236
View File
@@ -0,0 +1,236 @@
// 假执行器:模拟客户端核心(配对 → WSS → 心跳 → 任务执行 → 挑战处理)。
// 用途:D4 本地闭环验收 + 联调诊断;D8 起由 client/core 真 Agent 取代。
package main
import (
"bytes"
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/hex"
"encoding/json"
"flag"
"io"
"log"
"net/http"
"os"
"os/signal"
"path/filepath"
"strconv"
"syscall"
"time"
"github.com/coder/websocket"
"everypublish/shared/proto"
)
var (
server = flag.String("server", "http://127.0.0.1:8090", "服务器地址")
user = flag.String("user", "demo@everypublish.io", "登录邮箱")
password = flag.String("pass", "demo-pass-2026", "登录密码")
devName = flag.String("name", "fake-agent-01", "设备名")
execMS = flag.Int("exec", 50, "模拟执行时长 ms")
keyPath = flag.String("key", "./data/fake-agent.ed25519", "设备私钥文件")
)
func main() {
flag.Parse()
log.SetFlags(log.Ltime | log.Lmicroseconds)
access := login(*server, *user, *password)
code := pairCode(*server, access)
pub, priv := loadOrCreateKey(*keyPath)
deviceID := pair(*server, code, pub)
log.Printf("paired device %d", deviceID)
wsURL := "ws" + (*server)[4:] + "/ws/agent"
conn, _, err := websocket.Dial(context.Background(), wsURL, nil)
if err != nil {
log.Fatalf("dial: %v", err)
}
defer conn.Close(websocket.StatusNormalClosure, "bye")
nonce := randHex(16)
ts := time.Now().UnixMilli()
sig := ed25519.Sign(priv, []byte(nonce+"|"+strconv.FormatInt(ts, 10)))
hello := proto.NewEnvelope(randHex(8), proto.TypeHello, proto.DeviceHello{
DeviceID: strconv.FormatUint(deviceID, 10),
Nonce: nonce,
TS: ts,
Sig: base64.StdEncoding.EncodeToString(sig),
Version: "0.1.0-fake",
})
if err = send(conn, hello); err != nil {
log.Fatalf("hello: %v", err)
}
_, raw, err := conn.Read(context.Background())
if err != nil {
log.Fatalf("hello.ack read: %v", err)
}
var ackEnv proto.Envelope
_ = json.Unmarshal(raw, &ackEnv)
if ackEnv.Type != proto.TypeHelloAck {
log.Fatalf("expect hello.ack, got %s", ackEnv.Type)
}
var ack proto.HelloAck
_ = json.Unmarshal(ackEnv.Payload, &ack)
log.Printf("online (session %s)", ack.SessionID)
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stop()
go heartbeatLoop(ctx, conn, deviceID)
for {
_, raw, err = conn.Read(ctx)
if err != nil {
log.Fatalf("read: %v", err)
}
var env proto.Envelope
if err = json.Unmarshal(raw, &env); err != nil {
continue
}
switch env.Type {
case proto.TypeTaskPush:
var push proto.TaskPush
_ = json.Unmarshal(env.Payload, &push)
log.Printf("TASK push %s platform=%s title=%s", push.TaskID, push.Platform, push.Title)
_ = send(conn, proto.NewEnvelope(env.ID, proto.TypeTaskAck, proto.TaskAck{TaskID: push.TaskID, Accept: true}))
time.Sleep(time.Duration(*execMS) * time.Millisecond)
result := proto.TaskResult{
TaskID: push.TaskID,
Status: "success",
PublishedURL: "https://example.com/fake/" + push.TaskID,
Receipts: []string{},
FinishedAt: time.Now().UnixMilli(),
}
_ = send(conn, proto.NewEnvelope(randHex(8), proto.TypeTaskResult, result))
log.Printf("TASK %s done (%dms)", push.TaskID, *execMS)
case proto.TypeChallenge:
var ch proto.Challenge
_ = json.Unmarshal(env.Payload, &ch)
log.Printf("CHALLENGE %s kind=%s prompt=%s", ch.ChallengeID, ch.Kind, ch.Prompt)
_ = send(conn, proto.NewEnvelope(env.ID, proto.TypeChallengeAck, proto.ChallengeAck{ChallengeID: ch.ChallengeID, Action: "accept"}))
time.Sleep(time.Duration(*execMS) * time.Millisecond)
_ = send(conn, proto.NewEnvelope(randHex(8), proto.TypeChallengeSolve, proto.ChallengeSolve{ChallengeID: ch.ChallengeID, Value: "fake-qr-token"}))
log.Printf("CHALLENGE %s solved", ch.ChallengeID)
case proto.TypeHeartbeatAck:
default:
log.Printf("ignore msg type %s", env.Type)
}
}
}
func heartbeatLoop(ctx context.Context, conn *websocket.Conn, deviceID uint64) {
ticker := time.NewTicker(10 * time.Second)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
_ = send(conn, proto.NewEnvelope(randHex(8), proto.TypeHeartbeat, proto.Heartbeat{DeviceID: strconv.FormatUint(deviceID, 10), TS: time.Now().UnixMilli()}))
}
}
}
func send(conn *websocket.Conn, msg *proto.Envelope) error {
raw, err := json.Marshal(msg)
if err != nil {
return err
}
wctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
return conn.Write(wctx, websocket.MessageText, raw)
}
// ---- REST 辅助 ----
func httpJSON(method, url, token string, body interface{}) (int, []byte) {
var buf bytes.Buffer
if body != nil {
raw, _ := json.Marshal(body)
buf.Write(raw)
}
req, _ := http.NewRequest(method, url, &buf)
req.Header.Set("Content-Type", "application/json")
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
client := &http.Client{Timeout: 10 * time.Second}
resp, err := client.Do(req)
if err != nil {
log.Fatalf("http %s %s: %v", method, url, err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
return resp.StatusCode, raw
}
func field(data []byte, key string) string {
var m map[string]interface{}
_ = json.Unmarshal(data, &m)
if d, ok := m["data"].(map[string]interface{}); ok {
if v, ok := d[key]; ok {
switch t := v.(type) {
case string:
return t
case float64:
return strconv.FormatInt(int64(t), 10)
}
}
}
return ""
}
func login(base, email, password string) string {
code, raw := httpJSON("POST", base+"/api/v1/auth/login", "", ginH{"email": email, "password": password})
if code != 200 {
log.Fatalf("login failed: %s", string(raw))
}
log.Println("login ok")
return field(raw, "accessToken")
}
func pairCode(base, token string) string {
code, raw := httpJSON("POST", base+"/api/v1/agent/pair-code", token, nil)
if code != 200 {
log.Fatalf("pair-code failed: %s", string(raw))
}
return field(raw, "code")
}
func pair(base, code, pub string) uint64 {
httpCode, raw := httpJSON("POST", base+"/api/v1/agent/pair", "", ginH{"code": code, "deviceName": *devName, "os": "fake", "version": "0.1.0", "publicKey": pub})
if httpCode != 200 {
log.Fatalf("pair failed: %s", string(raw))
}
id, _ := strconv.ParseUint(field(raw, "deviceId"), 10, 64)
return id
}
type ginH map[string]interface{}
func loadOrCreateKey(path string) (string, ed25519.PrivateKey) {
if raw, err := os.ReadFile(path); err == nil {
seed, err := hex.DecodeString(string(raw))
if err == nil && len(seed) == ed25519.SeedSize {
priv := ed25519.NewKeyFromSeed(seed)
return base64.StdEncoding.EncodeToString(priv.Public().(ed25519.PublicKey)), priv
}
}
_, priv, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
log.Fatalf("keygen: %v", err)
}
_ = os.MkdirAll(filepath.Dir(path), 0o755)
_ = os.WriteFile(path, []byte(hex.EncodeToString(priv.Seed())), 0o600)
return base64.StdEncoding.EncodeToString(priv.Public().(ed25519.PublicKey)), priv
}
func randHex(n int) string {
b := make([]byte, n)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
+208
View File
@@ -0,0 +1,208 @@
-- EveryPublish 0001_init:初始 Schema(MySQL 8.x, utf8mb4)
-- 开发期由 GORM AutoMigrate 建表;生产执行本文件(golang-migrate)。
SET NAMES utf8mb4;
CREATE TABLE IF NOT EXISTS users (
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT,
created_at DATETIME(3) NULL,
updated_at DATETIME(3) NULL,
email VARCHAR(191) NOT NULL,
password_hash VARCHAR(255) NOT NULL,
nickname VARCHAR(64) NOT NULL DEFAULT '',
role VARCHAR(16) NOT NULL DEFAULT 'user',
totp_secret VARCHAR(64) NOT NULL DEFAULT '',
totp_enabled TINYINT(1) NOT NULL DEFAULT 0,
status VARCHAR(16) NOT NULL DEFAULT 'active',
PRIMARY KEY (id),
UNIQUE KEY uk_users_email (email)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS workspaces (
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT,
created_at DATETIME(3) NULL,
updated_at DATETIME(3) NULL,
name VARCHAR(128) NOT NULL DEFAULT '',
owner_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
plan VARCHAR(32) NOT NULL DEFAULT 'free',
status VARCHAR(16) NOT NULL DEFAULT 'active',
PRIMARY KEY (id),
KEY idx_workspaces_owner_id (owner_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS members (
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,
user_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
role VARCHAR(16) NOT NULL DEFAULT 'viewer',
PRIMARY KEY (id),
UNIQUE KEY uk_ws_user (workspace_id, user_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS accounts (
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,
platform VARCHAR(32) NOT NULL DEFAULT '',
account_name VARCHAR(128) NOT NULL DEFAULT '',
avatar_url VARCHAR(512) NOT NULL DEFAULT '',
status VARCHAR(16) NOT NULL DEFAULT 'unbound',
health INT NOT NULL DEFAULT 100,
agent_device_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
ip_profile VARCHAR(191) NOT NULL DEFAULT '',
last_active_at DATETIME(3) NULL,
last_error VARCHAR(512) NOT NULL DEFAULT '',
PRIMARY KEY (id),
KEY idx_accounts_workspace_id (workspace_id),
KEY idx_accounts_platform (platform),
KEY idx_accounts_agent_device_id (agent_device_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS materials (
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,
name VARCHAR(255) NOT NULL DEFAULT '',
kind VARCHAR(16) NOT NULL DEFAULT 'video',
size BIGINT NOT NULL DEFAULT 0,
sha256 VARCHAR(64) NOT NULL DEFAULT '',
mime VARCHAR(128) NOT NULL DEFAULT '',
storage_key VARCHAR(512) NOT NULL DEFAULT '',
`group` VARCHAR(64) NOT NULL DEFAULT '',
tags TEXT NULL,
status VARCHAR(16) NOT NULL DEFAULT 'ready',
created_by BIGINT UNSIGNED NOT NULL DEFAULT 0,
PRIMARY KEY (id),
KEY idx_materials_workspace_id (workspace_id),
KEY idx_materials_sha256 (sha256)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS tasks (
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,
title VARCHAR(255) NOT NULL DEFAULT '',
content TEXT NULL,
tags TEXT NULL,
account_ids TEXT NULL,
material_ids TEXT NULL,
schedule_at DATETIME(3) NULL,
status VARCHAR(32) NOT NULL DEFAULT 'draft',
priority INT NOT NULL DEFAULT 5,
retry_count INT NOT NULL DEFAULT 0,
max_retry INT NOT NULL DEFAULT 3,
error_message VARCHAR(1024) NOT NULL DEFAULT '',
published_at DATETIME(3) NULL,
published_urls TEXT NULL,
receipts TEXT NULL,
created_by BIGINT UNSIGNED NOT NULL DEFAULT 0,
reviewed_by BIGINT UNSIGNED NOT NULL DEFAULT 0,
review_note VARCHAR(512) NOT NULL DEFAULT '',
PRIMARY KEY (id),
KEY idx_tasks_workspace_id (workspace_id),
KEY idx_tasks_status (status),
KEY idx_tasks_schedule_at (schedule_at)
) 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,
updated_at DATETIME(3) NULL,
workspace_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
name VARCHAR(128) NOT NULL DEFAULT '',
os VARCHAR(32) NOT NULL DEFAULT '',
version VARCHAR(32) NOT NULL DEFAULT '',
public_key TEXT NULL,
status VARCHAR(16) NOT NULL DEFAULT 'offline',
ip VARCHAR(64) NOT NULL DEFAULT '',
revoked TINYINT(1) NOT NULL DEFAULT 0,
last_seen_at DATETIME(3) NULL,
paired_at DATETIME(3) NULL,
PRIMARY KEY (id),
KEY idx_agent_devices_workspace_id (workspace_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS challenges (
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,
account_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
device_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
platform VARCHAR(32) NOT NULL DEFAULT '',
kind VARCHAR(16) NOT NULL DEFAULT 'pending',
status VARCHAR(16) NOT NULL DEFAULT 'active',
qr_token VARCHAR(512) NOT NULL DEFAULT '',
qr_url VARCHAR(1024) NOT NULL DEFAULT '',
prompt VARCHAR(512) NOT NULL DEFAULT '',
payload TEXT NULL,
expires_at DATETIME(3) NULL,
solved_at DATETIME(3) NULL,
PRIMARY KEY (id),
KEY idx_challenges_workspace_id (workspace_id),
KEY idx_challenges_account_id (account_id),
KEY idx_challenges_device_id (device_id),
KEY idx_challenges_expires_at (expires_at)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS audit_logs (
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,
user_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
action VARCHAR(64) NOT NULL DEFAULT '',
resource VARCHAR(128) NOT NULL DEFAULT '',
detail TEXT NULL,
ip VARCHAR(64) NOT NULL DEFAULT '',
PRIMARY KEY (id),
KEY idx_audit_logs_workspace_id (workspace_id),
KEY idx_audit_logs_user_id (user_id),
KEY idx_audit_logs_action (action)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS notifications (
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,
user_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
kind VARCHAR(32) NOT NULL DEFAULT 'system',
title VARCHAR(255) NOT NULL DEFAULT '',
content TEXT NULL,
`read` TINYINT(1) NOT NULL DEFAULT 0,
PRIMARY KEY (id),
KEY idx_notifications_workspace_id (workspace_id),
KEY idx_notifications_user_id (user_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS pairing_codes (
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,
code VARCHAR(12) NOT NULL DEFAULT '',
expires_at DATETIME(3) NULL,
used TINYINT(1) NOT NULL DEFAULT 0,
device_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
PRIMARY KEY (id),
UNIQUE KEY uk_pairing_codes_code (code),
KEY idx_pairing_codes_workspace_id (workspace_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
CREATE TABLE IF NOT EXISTS transfer_tokens (
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT,
created_at DATETIME(3) NULL,
updated_at DATETIME(3) NULL,
material_id BIGINT UNSIGNED NOT NULL DEFAULT 0,
token VARCHAR(191) NOT NULL DEFAULT '',
expires_at DATETIME(3) NULL,
PRIMARY KEY (id),
UNIQUE KEY uk_transfer_tokens_token (token),
KEY idx_transfer_tokens_material_id (material_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
+35
View File
@@ -0,0 +1,35 @@
services:
mysql:
image: mysql:8.4
container_name: everypublish-mysql
environment:
MYSQL_ROOT_PASSWORD: everypublish
MYSQL_DATABASE: everypublish
TZ: Asia/Shanghai
command: --character-set-server=utf8mb4 --collation-server=utf8mb4_unicode_ci
ports:
- "3306:3306"
volumes:
- mysql_data:/var/lib/mysql
healthcheck:
test: ["CMD", "mysqladmin", "ping", "-h", "127.0.0.1", "-peverypublish"]
interval: 5s
timeout: 3s
retries: 30
redis:
image: redis:7-alpine
container_name: everypublish-redis
ports:
- "6379:6379"
volumes:
- redis_data:/data
healthcheck:
test: ["CMD", "redis-cli", "ping"]
interval: 5s
timeout: 3s
retries: 30
volumes:
mysql_data:
redis_data:
+55
View File
@@ -0,0 +1,55 @@
module everypublish/server
go 1.25.0
require (
everypublish/shared v0.0.0-00010101000000-000000000000
github.com/coder/websocket v1.8.15
github.com/gin-gonic/gin v1.12.0
github.com/golang-jwt/jwt/v5 v5.3.1
github.com/google/uuid v1.6.0
github.com/joho/godotenv v1.5.1
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/gorm v1.31.2
)
require (
filippo.io/edwards25519 v1.1.0 // indirect
github.com/bytedance/gopkg v0.1.3 // indirect
github.com/bytedance/sonic v1.15.0 // indirect
github.com/bytedance/sonic/loader v0.5.0 // indirect
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/cloudwego/base64x v0.1.6 // indirect
github.com/gabriel-vasile/mimetype v1.4.12 // indirect
github.com/gin-contrib/sse v1.1.0 // indirect
github.com/go-playground/locales v0.14.1 // indirect
github.com/go-playground/universal-translator v0.18.1 // indirect
github.com/go-playground/validator/v10 v10.30.1 // indirect
github.com/go-sql-driver/mysql v1.8.1 // indirect
github.com/goccy/go-json v0.10.5 // indirect
github.com/goccy/go-yaml v1.19.2 // indirect
github.com/jinzhu/inflection v1.0.0 // indirect
github.com/jinzhu/now v1.1.5 // indirect
github.com/json-iterator/go v1.1.12 // indirect
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/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
github.com/quic-go/qpack v0.6.0 // indirect
github.com/quic-go/quic-go v0.59.0 // indirect
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
github.com/ugorji/go/codec v1.3.1 // indirect
go.mongodb.org/mongo-driver/v2 v2.5.0 // indirect
go.uber.org/atomic v1.11.0 // indirect
golang.org/x/arch v0.22.0 // indirect
golang.org/x/net v0.57.0 // indirect
golang.org/x/sys v0.47.0 // indirect
golang.org/x/text v0.41.0 // indirect
google.golang.org/protobuf v1.36.10 // indirect
)
replace everypublish/shared => ../shared
+125
View File
@@ -0,0 +1,125 @@
filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA=
filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4=
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M=
github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM=
github.com/bytedance/sonic v1.15.0 h1:/PXeWFaR5ElNcVE84U0dOHjiMHQOwNIx3K4ymzh/uSE=
github.com/bytedance/sonic v1.15.0/go.mod h1:tFkWrPz0/CUCLEF4ri4UkHekCIcdnkqXw9VduqpJh0k=
github.com/bytedance/sonic/loader v0.5.0 h1:gXH3KVnatgY7loH5/TkeVyXPfESoqSBSBEiDd5VjlgE=
github.com/bytedance/sonic/loader v0.5.0/go.mod h1:AR4NYCk5DdzZizZ5djGqQ92eEhCCcdf5x77udYiSJRo=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/cloudwego/base64x v0.1.6 h1:t11wG9AECkCDk5fMSoxmufanudBtJ+/HemLstXDLI2M=
github.com/cloudwego/base64x v0.1.6/go.mod h1:OFcloc187FXDaYHvrNIjxSe8ncn0OOM8gEHfghB2IPU=
github.com/coder/websocket v1.8.15 h1:6B2JPeOGlpff2Uz6vOEH1Vzpi0iUz20A+lPVhPHtNUA=
github.com/coder/websocket v1.8.15/go.mod h1:NX3SzP+inril6yawo5CQXx8+fk145lPDC6pumgx0mVg=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/gabriel-vasile/mimetype v1.4.12 h1:e9hWvmLYvtp846tLHam2o++qitpguFiYCKbn0w9jyqw=
github.com/gabriel-vasile/mimetype v1.4.12/go.mod h1:d+9Oxyo1wTzWdyVUPMmXFvp4F9tea18J8ufA774AB3s=
github.com/gin-contrib/sse v1.1.0 h1:n0w2GMuUpWDVp7qSpvze6fAu9iRxJY4Hmj6AmBOU05w=
github.com/gin-contrib/sse v1.1.0/go.mod h1:hxRZ5gVpWMT7Z0B0gSNYqqsSCNIJMjzvm6fqCz9vjwM=
github.com/gin-gonic/gin v1.12.0 h1:b3YAbrZtnf8N//yjKeU2+MQsh2mY5htkZidOM7O0wG8=
github.com/gin-gonic/gin v1.12.0/go.mod h1:VxccKfsSllpKshkBWgVgRniFFAzFb9csfngsqANjnLc=
github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s=
github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4=
github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/oXslEjJA=
github.com/go-playground/locales v0.14.1/go.mod h1:hxrqLVvrK65+Rwrd5Fc6F2O76J/NuW9t0sjnWqG1slY=
github.com/go-playground/universal-translator v0.18.1 h1:Bcnm0ZwsGyWbCzImXv+pAJnYK9S473LQFuzCbDbfSFY=
github.com/go-playground/universal-translator v0.18.1/go.mod h1:xekY+UJKNuX9WP91TpwSH2VMlDf28Uj24BCp08ZFTUY=
github.com/go-playground/validator/v10 v10.30.1 h1:f3zDSN/zOma+w6+1Wswgd9fLkdwy06ntQJp0BBvFG0w=
github.com/go-playground/validator/v10 v10.30.1/go.mod h1:oSuBIQzuJxL//3MelwSLD5hc2Tu889bF0Idm9Dg26cM=
github.com/go-sql-driver/mysql v1.8.1 h1:LedoTUt/eveggdHS9qUFC1EFSa8bU2+1pZjSRpvNJ1Y=
github.com/go-sql-driver/mysql v1.8.1/go.mod h1:wEBSXgmK//2ZFJyE+qWnIsVGmvmEKlqwuVSjsCm7DZg=
github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4=
github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M=
github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM=
github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA=
github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY=
github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E=
github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc=
github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8=
github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0=
github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4=
github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM=
github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo=
github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y=
github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ=
github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI=
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU=
github.com/mattn/go-sqlite3 v1.14.22/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y=
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg=
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M=
github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk=
github.com/pelletier/go-toml/v2 v2.2.4 h1:mye9XuhQ6gvn5h28+VilKrrPoQVanw5PMw/TB0t5Ec4=
github.com/pelletier/go-toml/v2 v2.2.4/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8=
github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII=
github.com/quic-go/quic-go v0.59.0 h1:OLJkp1Mlm/aS7dpKgTc6cnpynnD2Xg7C1pwL6vy/SAw=
github.com/quic-go/quic-go v0.59.0/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU=
github.com/redis/go-redis/v9 v9.22.0 h1:laDvpYXTJtZLloinw1fA5Kqd6HAEH2XKxOkG/PDq2F0=
github.com/redis/go-redis/v9 v9.22.0/go.mod h1:y2g0Wj8rQvuK0ELM+oxSudcLtC09JScs98I/X9gRWY4=
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw=
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI=
github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08=
github.com/ugorji/go/codec v1.3.1 h1:waO7eEiFDwidsBN6agj1vJQ4AG7lh2yqXyOXqhgQuyY=
github.com/ugorji/go/codec v1.3.1/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4=
github.com/zeebo/xxh3 v1.1.0 h1:s7DLGDK45Dyfg7++yxI0khrfwq9661w9EN78eP/UZVs=
github.com/zeebo/xxh3 v1.1.0/go.mod h1:IisAie1LELR4xhVinxWS5+zf1lA4p0MW4T+w+W07F5s=
go.mongodb.org/mongo-driver/v2 v2.5.0 h1:yXUhImUjjAInNcpTcAlPHiT7bIXhshCTL3jVBkF3xaE=
go.mongodb.org/mongo-driver/v2 v2.5.0/go.mod h1:yOI9kBsufol30iFsl1slpdq1I0eHPzybRWdyYUs8K/0=
go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE=
go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0=
go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y=
go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU=
golang.org/x/arch v0.22.0 h1:c/Zle32i5ttqRXjdLyyHZESLD/bB90DCU1g9l/0YBDI=
golang.org/x/arch v0.22.0/go.mod h1:dNHoOeKiyja7GTvF9NJS1l3Z2yntpQNzgrjh1cU103A=
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
google.golang.org/protobuf v1.36.10 h1:AYd7cD/uASjIL6Q9LiTjz8JLcrh/88q5UObnmY3aOOE=
google.golang.org/protobuf v1.36.10/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gorm.io/driver/mysql v1.6.0 h1:eNbLmNTpPpTOVZi8MMxCi2aaIm0ZpInbORNXDwyLGvg=
gorm.io/driver/mysql v1.6.0/go.mod h1:D/oCC2GWK3M/dqoLxnOlaNKmXz8WNTfcS9y5ovaSqKo=
gorm.io/driver/sqlite v1.6.0 h1:WHRRrIiulaPiPFmDcod6prc4l2VGVWHz80KspNsxSfQ=
gorm.io/driver/sqlite v1.6.0/go.mod h1:AO9V1qIQddBESngQUKWL9yoH93HIeA1X6V633rBwyT8=
gorm.io/gorm v1.31.2 h1:3o8FXNo9v9S858gil+3LlZA1LkCOzgb4g5BL64FgaCo=
gorm.io/gorm v1.31.2/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs=
+221
View File
@@ -0,0 +1,221 @@
package api_test
import (
"bytes"
"encoding/json"
"fmt"
"mime/multipart"
"net/http/httptest"
"net/textproto"
"strings"
"testing"
"github.com/gin-gonic/gin"
"everypublish/server/internal/models"
)
func multipartBody(t *testing.T, filename, contentType string, content []byte) *bytes.Buffer {
t.Helper()
buf := &bytes.Buffer{}
w := multipart.NewWriter(buf)
h := make(textproto.MIMEHeader)
h.Set("Content-Disposition", fmt.Sprintf(`form-data; name="file"; filename="%s"`, filename))
h.Set("Content-Type", contentType)
part, err := w.CreatePart(h)
if err != nil {
t.Fatalf("create part: %v", err)
}
_, _ = part.Write(content)
_ = w.Close()
return buf
}
func uploadMaterial(t *testing.T, access string, filename string, content []byte) (uint64, string) {
t.Helper()
buf := multipartBody(t, filename, "video/mp4", content)
req := httptest.NewRequest("POST", "/api/v1/materials", buf)
req.Header.Set("Content-Type", "multipart/form-data; boundary="+strings.Split(buf.String(), "\r\n")[0][2:])
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
testRouter.ServeHTTP(w, req)
var env envelope
_ = json.Unmarshal(w.Body.Bytes(), &env)
if w.Code != 200 || env.Code != 0 {
t.Fatalf("upload failed: %d %s", w.Code, env.Message)
}
var data struct {
Material models.Material `json:"material"`
Dedup bool `json:"dedup"`
}
_ = json.Unmarshal(env.Data, &data)
if data.Dedup {
t.Fatal("first upload should not be dedup")
}
return data.Material.ID, data.Material.SHA256
}
func TestAccountChallengeFlow(t *testing.T) {
resetTestDB()
access, _ := register(t, "ops@test.com")
code, env := doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "douyin", "accountName": "抖音测试号"})
if code != 200 || env.Code != 0 {
t.Fatalf("create account failed: %d %s", code, env.Message)
}
var acc models.Account
_ = json.Unmarshal(env.Data, &acc)
// 不支持的平台
code, _ = doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "tiktok", "accountName": "x"})
if code != 400 {
t.Fatalf("unsupported platform should 400, got %d", code)
}
// 绑定 → 挑战
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/accounts/%d/bind", acc.ID), access, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("bind failed: %d %s", code, env.Message)
}
// 绑定中账号不可删
code, _ = doReq(t, "DELETE", fmt.Sprintf("/api/v1/accounts/%d", acc.ID), access, nil)
if code != 409 {
t.Fatalf("delete binding account should 409, got %d", code)
}
// 挑战列表
code, env = doReq(t, "GET", "/api/v1/challenges", access, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("challenge list failed: %d %s", code, env.Message)
}
var cl struct {
List []models.Challenge `json:"list"`
}
_ = json.Unmarshal(env.Data, &cl)
if len(cl.List) != 1 || cl.List[0].Kind != "pending" {
t.Fatalf("expect 1 pending challenge, got %d", len(cl.List))
}
// 挂起 → 重发
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/challenges/%d/suspend", cl.List[0].ID), access, nil)
if code != 200 {
t.Fatalf("suspend failed: %d", code)
}
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/challenges/%d/resend", cl.List[0].ID), access, nil)
if code != 200 {
t.Fatalf("resend failed: %d", code)
}
}
func TestMaterialAndTaskFlow(t *testing.T) {
resetTestDB()
access, _ := register(t, "pub@test.com")
content := []byte("fake-video-bytes-for-sha256-test-0001")
matID, sha := uploadMaterial(t, access, "demo.mp4", content)
if len(sha) != 64 {
t.Fatalf("bad sha256: %s", sha)
}
// 重复上传 → 去重
buf := multipartBody(t, "demo2.mp4", "video/mp4", content)
req := httptest.NewRequest("POST", "/api/v1/materials", buf)
req.Header.Set("Content-Type", "multipart/form-data; boundary="+strings.Split(buf.String(), "\r\n")[0][2:])
req.Header.Set("Authorization", "Bearer "+access)
w := httptest.NewRecorder()
testRouter.ServeHTTP(w, req)
var env envelope
_ = json.Unmarshal(w.Body.Bytes(), &env)
var data struct {
Material models.Material `json:"material"`
Dedup bool `json:"dedup"`
}
_ = json.Unmarshal(env.Data, &data)
if !data.Dedup || data.Material.ID != matID {
t.Fatalf("dedup expect same id %d, got dedup=%v id=%d", matID, data.Dedup, data.Material.ID)
}
// 签名直链(一次性)
code, env := doReq(t, "GET", fmt.Sprintf("/api/v1/materials/%d/url", matID), access, nil)
if code != 200 {
t.Fatalf("url failed: %d", code)
}
var urld struct {
URL string `json:"url"`
}
_ = json.Unmarshal(env.Data, &urld)
path := strings.TrimPrefix(urld.URL, "http://127.0.0.1:8090")
if path == urld.URL {
t.Fatalf("bad url: %s", urld.URL)
}
code, _ = doReq(t, "GET", path, "", nil)
if code != 200 {
t.Fatalf("file fetch should 200, got %d", code)
}
code, _ = doReq(t, "GET", path, "", nil)
if code != 404 {
t.Fatalf("one-time token should 404 on reuse, got %d", code)
}
// 建任务 → 提交 → 驳回 → 重提 → 通过 → 取消
code, env = doReq(t, "POST", "/api/v1/tasks", access, gin.H{
"title": "新品发布", "content": "正文", "tags": []string{"新品"},
"accountIds": []uint64{1}, "materialIds": []uint64{matID},
})
if code != 200 || env.Code != 0 {
t.Fatalf("create task failed: %d %s", code, env.Message)
}
var taskRow models.Task
_ = json.Unmarshal(env.Data, &taskRow)
if taskRow.Status != "draft" {
t.Fatalf("new task should be draft, got %s", taskRow.Status)
}
// 非法转移:draft 直接 approve 应 409
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/approve", taskRow.ID), access, nil)
if code != 409 {
t.Fatalf("draft approve should 409, got %d", code)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/submit", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("submit failed: %d", code)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/reject", taskRow.ID), access, gin.H{"note": "标题需修改"})
if code != 200 {
t.Fatalf("reject failed: %d", code)
}
code, env = doReq(t, "GET", fmt.Sprintf("/api/v1/tasks/%d", taskRow.ID), access, nil)
_ = json.Unmarshal(env.Data, &taskRow)
if taskRow.Status != "rejected" {
t.Fatalf("expect rejected, got %s", taskRow.Status)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/resubmit", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("resubmit failed: %d", code)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/approve", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("approve failed: %d", code)
}
code, env = doReq(t, "GET", fmt.Sprintf("/api/v1/tasks/%d", taskRow.ID), access, nil)
_ = json.Unmarshal(env.Data, &taskRow)
if taskRow.Status != "queued" {
t.Fatalf("expect queued, got %s", taskRow.Status)
}
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/cancel", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("cancel failed: %d", code)
}
// 通知:驳回+通过 2 条
code, env = doReq(t, "GET", "/api/v1/notifications", access, nil)
if code != 200 {
t.Fatalf("notifications failed: %d", code)
}
var nl struct {
List []models.Notification `json:"list"`
Unread int64 `json:"unread"`
}
_ = json.Unmarshal(env.Data, &nl)
if nl.Unread < 2 {
t.Fatalf("expect >=2 unread, got %d", nl.Unread)
}
code, _ = doReq(t, "POST", "/api/v1/notifications/read-all", access, nil)
if code != 200 {
t.Fatalf("read-all failed: %d", code)
}
code, env = doReq(t, "GET", "/api/v1/notifications", access, nil)
_ = json.Unmarshal(env.Data, &nl)
if nl.Unread != 0 {
t.Fatalf("expect 0 unread after read-all, got %d", nl.Unread)
}
}
+195
View File
@@ -0,0 +1,195 @@
package handlers
import (
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
"everypublish/server/internal/ws"
"everypublish/shared/proto"
)
// AccountHandler 账号台账处理器(凭据仅在 Agent 本机保险库,服务器零凭据)
type AccountHandler struct {
DB *gorm.DB
Hub *ws.Hub
}
var platforms = map[string]bool{
"douyin": true, "kuaishou": true, "xiaohongshu": true, "bilibili": true,
"shipinhao": true, "x": true, "instagram": true, "whatsapp": true, "youtube": true,
}
type accountReq struct {
Platform string `json:"platform" binding:"required"`
Remark string `json:"remark"`
AgentDeviceID uint64 `json:"agentDeviceId"`
AccountName string `json:"accountName"`
}
// List 账号台账(筛选 platform/status)
func (h *AccountHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.Account{}).Where("workspace_id = ?", c.GetUint64("wsid"))
if p := c.Query("platform"); p != "" {
q = q.Where("platform = ?", p)
}
if s := c.Query("status"); s != "" {
q = q.Where("status = ?", s)
}
var total int64
q.Count(&total)
var list []models.Account
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&list)
response.OK(c, gin.H{"list": list, "total": total, "page": page, "size": size})
}
// Create 新增账号(初始 unbound;凭据只登记元数据)
func (h *AccountHandler) Create(c *gin.Context) {
var req accountReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
if !platforms[req.Platform] {
response.Fail(c, http.StatusBadRequest, 1001, "暂不支持该平台")
return
}
account := models.Account{
WorkspaceID: c.GetUint64("wsid"),
Platform: req.Platform,
Remark: req.Remark,
AgentDeviceID: req.AgentDeviceID,
Status: "unbound",
Health: 100,
}
if err := h.DB.Create(&account).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "创建失败")
return
}
response.Audit(c, "account.create", "account:"+itoa(account.ID), gin.H{"platform": req.Platform, "remark": req.Remark})
response.OK(c, account)
}
// Update 修改账号资料(名称/头像/IP画像)
func (h *AccountHandler) Update(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req accountReq
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
updates := map[string]interface{}{"remark": req.Remark}
if req.AccountName != "" {
updates["account_name"] = req.AccountName
}
if req.AgentDeviceID > 0 {
updates["agent_device_id"] = req.AgentDeviceID
}
if err = h.DB.Model(&models.Account{}).Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).Updates(updates).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "account.update", "account:"+itoa(id), gin.H{"remark": req.Remark})
response.OK(c, nil)
}
// Delete 删除账号(仅未绑定)
func (h *AccountHandler) Delete(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var account models.Account
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&account).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "账号不存在")
return
}
if account.Status == "active" || account.Status == "binding" {
response.Fail(c, http.StatusConflict, 3001, "已绑定或绑定中的账号不可删除")
return
}
if err = h.DB.Delete(&account).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "删除失败")
return
}
response.Audit(c, "account.delete", "account:"+itoa(id), nil)
response.OK(c, nil)
}
// Bind 发起绑定:创建挑战记录(pending),等待客户端扫码(D4 WSS 联动)
func (h *AccountHandler) Bind(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var account models.Account
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&account).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "账号不存在")
return
}
if account.Status == "active" {
response.Fail(c, http.StatusConflict, 3002, "账号已绑定")
return
}
challenge := models.Challenge{
WorkspaceID: c.GetUint64("wsid"),
AccountID: account.ID,
Platform: account.Platform,
Kind: "pending",
Status: "active",
Prompt: "等待客户端发起扫码,稍后此处展示二维码",
ExpiresAt: time.Now().Add(30 * time.Minute),
}
err = h.DB.Transaction(func(tx *gorm.DB) error {
if err = tx.Create(&challenge).Error; err != nil {
return err
}
return tx.Model(&account).Update("status", "binding").Error
})
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "发起绑定失败")
return
}
// 定向下发:优先账号指定设备,否则工作空间首个在线设备
if h.Hub != nil {
target := account.AgentDeviceID
if target == 0 || !h.Hub.AgentOnline(target) {
if online := h.Hub.OnlineDevices(c.GetUint64("wsid")); len(online) > 0 {
target = online[0]
}
}
if target != 0 {
env := proto.NewEnvelope(uuid.NewString(), proto.TypeChallenge, proto.Challenge{
ChallengeID: strconv.FormatUint(challenge.ID, 10),
AccountID: strconv.FormatUint(account.ID, 10),
Platform: account.Platform,
Kind: "qr",
Prompt: "请使用手机客户端扫码登录 " + account.Platform,
ExpiresAt: challenge.ExpiresAt.UnixMilli(),
})
_ = h.Hub.SendToAgent(target, env)
}
}
response.Audit(c, "account.bind", "account:"+itoa(account.ID), gin.H{"challengeId": challenge.ID})
response.OK(c, gin.H{"challengeId": challenge.ID, "account": account.ID, "status": "binding"})
}
+143
View File
@@ -0,0 +1,143 @@
package handlers
import (
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
"everypublish/server/internal/ws"
)
// AdminHandler 平台管理处理器(AdminRequired 守卫)
type AdminHandler struct {
DB *gorm.DB
Hub *ws.Hub
}
type tenantRow struct {
models.Workspace
OwnerEmail string `json:"ownerEmail"`
MemberCount int64 `json:"memberCount"`
DeviceCount int64 `json:"deviceCount"`
}
// Tenants 租户列表(含 owner 邮箱与成员/设备数)
func (h *AdminHandler) Tenants(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.Workspace{})
if s := c.Query("status"); s != "" {
q = q.Where("status = ?", s)
}
var total int64
q.Count(&total)
var wsList []models.Workspace
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&wsList)
rows := make([]tenantRow, 0, len(wsList))
for _, ws := range wsList {
row := tenantRow{Workspace: ws}
var owner models.User
if err := h.DB.First(&owner, ws.OwnerID).Error; err == nil {
row.OwnerEmail = owner.Email
}
h.DB.Model(&models.Member{}).Where("workspace_id = ?", ws.ID).Count(&row.MemberCount)
h.DB.Model(&models.AgentDevice{}).Where("workspace_id = ?", ws.ID).Count(&row.DeviceCount)
rows = append(rows, row)
}
response.OK(c, gin.H{"list": rows, "total": total, "page": page, "size": size})
}
// TenantStatus 启用/停用租户
func (h *AdminHandler) TenantStatus(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req struct {
Status string `json:"status" binding:"required"`
}
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
if req.Status != "active" && req.Status != "suspended" {
response.Fail(c, http.StatusBadRequest, 1001, "状态不合法")
return
}
var ws models.Workspace
if err = h.DB.First(&ws, id).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "租户不存在")
return
}
if err = h.DB.Model(&ws).Update("status", req.Status).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "admin.tenant.status", "workspace:"+itoa(id), gin.H{"status": req.Status})
response.OK(c, nil)
}
type agentRow struct {
models.AgentDevice
WorkspaceName string `json:"workspaceName"`
}
// Agents 全局 Agent 总览
func (h *AdminHandler) Agents(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
var total int64
h.DB.Model(&models.AgentDevice{}).Count(&total)
var devices []models.AgentDevice
h.DB.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&devices)
rows := make([]agentRow, 0, len(devices))
for _, d := range devices {
row := agentRow{AgentDevice: d}
var ws models.Workspace
if err := h.DB.First(&ws, d.WorkspaceID).Error; err == nil {
row.WorkspaceName = ws.Name
}
rows = append(rows, row)
}
response.OK(c, gin.H{"list": rows, "total": total, "page": page, "size": size})
}
// AgentRevoke 全局强制吊销
func (h *AdminHandler) AgentRevoke(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var device models.AgentDevice
if err = h.DB.First(&device, id).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "设备不存在")
return
}
if err = h.DB.Model(&device).Updates(map[string]interface{}{"revoked": true, "status": "offline"}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "吊销失败")
return
}
if h.Hub != nil {
h.Hub.KickAgent(device.ID)
}
response.Audit(c, "admin.agent.revoke", "device:"+itoa(id), nil)
response.OK(c, nil)
}
+147
View File
@@ -0,0 +1,147 @@
package handlers
import (
"crypto/rand"
"errors"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
"everypublish/server/internal/ws"
)
// AgentHandler 设备配对与列表
type AgentHandler struct {
DB *gorm.DB
Hub *ws.Hub
}
const pairAlphabet = "ABCDEFGHJKMNPQRSTUVWXYZ23456789"
var errPairCodeUsed = errors.New("pair code already used")
func randCode(n int) string {
b := make([]byte, n)
_, _ = rand.Read(b)
for i := range b {
b[i] = pairAlphabet[int(b[i])%len(pairAlphabet)]
}
return string(b)
}
// PairCode 生成配对码(5 分钟一次性)
func (h *AgentHandler) PairCode(c *gin.Context) {
pc := models.PairingCode{
WorkspaceID: c.GetUint64("wsid"),
Code: randCode(6),
ExpiresAt: time.Now().Add(5 * time.Minute),
}
if err := h.DB.Create(&pc).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "生成配对码失败")
return
}
response.Audit(c, "agent.pair_code", "pairing:"+itoa(pc.ID), gin.H{"code": pc.Code})
response.OK(c, gin.H{"code": pc.Code, "expiresAt": pc.ExpiresAt})
}
// Pair 设备配对(无鉴权,凭一次性码;注册 Ed25519 公钥)
func (h *AgentHandler) Pair(c *gin.Context) {
var req struct {
Code string `json:"code" binding:"required"`
DeviceName string `json:"deviceName" binding:"required"`
OS string `json:"os"`
Version string `json:"version"`
PublicKey string `json:"publicKey" binding:"required"`
}
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var pc models.PairingCode
if err := h.DB.Where("code = ? and used = ?", req.Code, false).First(&pc).Error; err != nil {
response.Fail(c, http.StatusBadRequest, 2006, "配对码无效")
return
}
if time.Now().After(pc.ExpiresAt) {
response.Fail(c, http.StatusGone, 2006, "配对码已过期")
return
}
device := models.AgentDevice{
WorkspaceID: pc.WorkspaceID,
Name: req.DeviceName,
OS: req.OS,
Version: req.Version,
PublicKey: req.PublicKey,
Status: "offline",
PairedAt: time.Now(),
}
err := h.DB.Transaction(func(tx *gorm.DB) error {
// 原子占用:仅当仍 unused 且未过期时置 used=true;RowsAffected==1 才视为抢到,
// 防止并发用同一码配出多台设备。
res := tx.Model(&models.PairingCode{}).
Where("id = ? and used = ?", pc.ID, false).
Update("used", true)
if res.Error != nil {
return res.Error
}
if res.RowsAffected != 1 {
return errPairCodeUsed
}
if err := tx.Create(&device).Error; err != nil {
return err
}
if err := tx.Model(&models.PairingCode{}).Where("id = ?", pc.ID).
Update("device_id", device.ID).Error; err != nil {
return err
}
return nil
})
if err != nil {
if err == errPairCodeUsed {
response.Fail(c, http.StatusBadRequest, 2006, "配对码已被使用")
return
}
response.Fail(c, http.StatusInternalServerError, 5000, "配对失败")
return
}
response.OK(c, gin.H{"deviceId": device.ID, "workspaceId": device.WorkspaceID})
}
// Devices 设备列表(在线状态来自 Hub,此处以 DB 状态兜底)
func (h *AgentHandler) Devices(c *gin.Context) {
var list []models.AgentDevice
h.DB.Where("workspace_id = ?", c.GetUint64("wsid")).Order("id desc").Find(&list)
response.OK(c, gin.H{"list": list})
}
// Revoke 吊销设备并强制断开(owner/admin)
func (h *AgentHandler) Revoke(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可吊销设备")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var device models.AgentDevice
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&device).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "设备不存在")
return
}
if err = h.DB.Model(&device).Updates(map[string]interface{}{"revoked": true, "status": "offline"}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "吊销失败")
return
}
if h.Hub != nil {
h.Hub.KickAgent(device.ID)
}
response.Audit(c, "agent.revoke", "device:"+itoa(id), nil)
response.OK(c, nil)
}
+40
View File
@@ -0,0 +1,40 @@
package handlers
import (
"strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// AuditHandler 审计日志处理器
type AuditHandler struct {
DB *gorm.DB
}
// List 审计日志(分页 + 筛选)
func (h *AuditHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.AuditLog{}).Where("workspace_id = ?", c.GetUint64("wsid"))
if uid := c.Query("userId"); uid != "" {
q = q.Where("user_id = ?", uid)
}
if action := c.Query("action"); action != "" {
q = q.Where("action = ?", action)
}
var total int64
q.Count(&total)
var list []models.AuditLog
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&list)
response.OK(c, gin.H{"list": list, "total": total, "page": page, "size": size})
}
+235
View File
@@ -0,0 +1,235 @@
package handlers
import (
"net/http"
"strings"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/auth"
"everypublish/server/internal/cache"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
)
// AuthHandler 认证处理器
type AuthHandler struct {
DB *gorm.DB
Rds *cache.Redis
Cfg *config.Config
}
type registerReq struct {
Email string `json:"email" binding:"required"`
Password string `json:"password" binding:"required,min=8"`
Nickname string `json:"nickname"`
}
type loginReq struct {
Email string `json:"email" binding:"required"`
Password string `json:"password" binding:"required"`
TotpCode string `json:"totpCode"`
}
type refreshReq struct {
RefreshToken string `json:"refreshToken" binding:"required"`
}
type logoutReq struct {
RefreshToken string `json:"refreshToken"`
}
// Register 注册:创建用户 + 默认工作空间 + owner 成员
func (h *AuthHandler) Register(c *gin.Context) {
var req registerReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法:"+err.Error())
return
}
req.Email = strings.ToLower(strings.TrimSpace(req.Email))
if !strings.Contains(req.Email, "@") {
response.Fail(c, http.StatusBadRequest, 1001, "邮箱格式不正确")
return
}
if req.Nickname == "" {
req.Nickname = strings.Split(req.Email, "@")[0]
}
var count int64
h.DB.Model(&models.User{}).Where("email = ?", req.Email).Count(&count)
if count > 0 {
response.Fail(c, http.StatusConflict, 2001, "邮箱已注册")
return
}
hash, err := auth.HashPassword(req.Password)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "服务内部错误")
return
}
err = h.DB.Transaction(func(tx *gorm.DB) error {
user := models.User{Email: req.Email, PasswordHash: hash, Nickname: req.Nickname, Role: "user", Status: "active"}
if err := tx.Create(&user).Error; err != nil {
return err
}
ws := models.Workspace{Name: req.Nickname + "的工作空间", OwnerID: user.ID, Plan: "free", Status: "active"}
if err := tx.Create(&ws).Error; err != nil {
return err
}
member := models.Member{WorkspaceID: ws.ID, UserID: user.ID, Role: "owner"}
if err := tx.Create(&member).Error; err != nil {
return err
}
c.Set("regUID", user.ID)
c.Set("regWSID", ws.ID)
return nil
})
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "注册失败")
return
}
uid, wsid := c.GetUint64("regUID"), c.GetUint64("regWSID")
access, refresh, err := h.issuePair(uid, wsid, "user", "owner")
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "签发令牌失败")
return
}
var user models.User
h.DB.First(&user, uid)
response.OK(c, gin.H{
"accessToken": access,
"refreshToken": refresh,
"user": userDTO(user),
"workspace": gin.H{"id": wsid, "name": req.Nickname + "的工作空间"},
})
}
// Login 登录(2FA 默认关;开启后需 totpCode)
func (h *AuthHandler) Login(c *gin.Context) {
var req loginReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var user models.User
if err := h.DB.Where("email = ?", strings.ToLower(strings.TrimSpace(req.Email))).First(&user).Error; err != nil {
response.Fail(c, http.StatusUnauthorized, 2002, "邮箱或密码错误")
return
}
ok, err := auth.VerifyPassword(req.Password, user.PasswordHash)
if err != nil || !ok {
response.Fail(c, http.StatusUnauthorized, 2002, "邮箱或密码错误")
return
}
if user.TOTPEnabled {
// 二期接入 TOTP 校验;一期默认关闭
response.Fail(c, http.StatusForbidden, 1006, "需要两步验证码")
return
}
if user.Status != "active" {
response.Fail(c, http.StatusForbidden, 2003, "账号已停用")
return
}
var member models.Member
if err := h.DB.Where("user_id = ?", user.ID).Order("id asc").First(&member).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "工作空间数据缺失")
return
}
access, refresh, err := h.issuePair(user.ID, member.WorkspaceID, user.Role, member.Role)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "签发令牌失败")
return
}
c.Set("wsid", member.WorkspaceID)
response.Audit(c, "auth.login", "user:"+itoa(user.ID), gin.H{"email": user.Email})
response.OK(c, gin.H{
"accessToken": access,
"refreshToken": refresh,
"user": userDTO(user),
"workspace": gin.H{"id": member.WorkspaceID},
})
}
// Refresh 刷新令牌轮换:旧 jti 吊销,签发新对
func (h *AuthHandler) Refresh(c *gin.Context) {
var req refreshReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
claims, err := auth.ParseRefresh(h.Cfg.JWTSecret, req.RefreshToken)
if err != nil {
response.Fail(c, http.StatusUnauthorized, 2003, "刷新令牌无效或已过期")
return
}
valid, err := h.Rds.ExistsRefresh(claims.UID, claims.JTI)
if err != nil || !valid {
response.Fail(c, http.StatusUnauthorized, 2003, "刷新令牌已失效")
return
}
if err = h.Rds.DeleteRefresh(claims.UID, claims.JTI); err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "服务内部错误")
return
}
var user models.User
if err = h.DB.First(&user, claims.UID).Error; err != nil {
response.Fail(c, http.StatusUnauthorized, 2003, "用户不存在")
return
}
var member models.Member
if err = h.DB.Where("user_id = ?", user.ID).Order("id asc").First(&member).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "工作空间数据缺失")
return
}
access, refresh, err := h.issuePair(user.ID, member.WorkspaceID, user.Role, member.Role)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "签发令牌失败")
return
}
response.OK(c, gin.H{"accessToken": access, "refreshToken": refresh})
}
// Logout 登出:吊销刷新令牌
func (h *AuthHandler) Logout(c *gin.Context) {
var req logoutReq
_ = c.ShouldBindJSON(&req)
if req.RefreshToken != "" {
if claims, err := auth.ParseRefresh(h.Cfg.JWTSecret, req.RefreshToken); err == nil {
_ = h.Rds.DeleteRefresh(claims.UID, claims.JTI)
}
}
response.Audit(c, "auth.logout", "user:"+itoa(c.GetUint64("uid")), nil)
response.OK(c, nil)
}
// Me 当前用户信息
func (h *AuthHandler) Me(c *gin.Context) {
var user models.User
if err := h.DB.First(&user, c.GetUint64("uid")).Error; err != nil {
response.Fail(c, http.StatusUnauthorized, 1002, "用户不存在")
return
}
response.OK(c, gin.H{"user": userDTO(user), "workspaceId": c.GetUint64("wsid"), "memberRole": c.GetString("mrole")})
}
// issuePair 签发令牌对并把 refresh jti 登记到 Redis
func (h *AuthHandler) issuePair(uid, wsid uint64, role, mrole string) (string, string, error) {
access, err := auth.IssueAccess(h.Cfg.JWTSecret, uid, wsid, role, mrole)
if err != nil {
return "", "", err
}
jti := uuid.NewString()
refresh, err := auth.IssueRefresh(h.Cfg.JWTSecret, uid, jti)
if err != nil {
return "", "", err
}
if err = h.Rds.SaveRefresh(uid, jti, auth.RefreshTTL); err != nil {
return "", "", err
}
return access, refresh, nil
}
func userDTO(u models.User) gin.H {
return gin.H{"id": u.ID, "email": u.Email, "nickname": u.Nickname, "role": u.Role, "totpEnabled": u.TOTPEnabled}
}
+114
View File
@@ -0,0 +1,114 @@
package handlers
import (
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// ChallengeHandler 二次验证挑战处理器
type ChallengeHandler struct {
DB *gorm.DB
}
// List 挑战列表(过期软置 expired)
func (h *ChallengeHandler) List(c *gin.Context) {
h.DB.Model(&models.Challenge{}).
Where("workspace_id = ? and status = ? and expires_at < ?", c.GetUint64("wsid"), "active", time.Now()).
Update("status", "expired")
q := h.DB.Model(&models.Challenge{}).Where("workspace_id = ?", c.GetUint64("wsid"))
if s := c.Query("status"); s != "" {
q = q.Where("status = ?", s)
}
if a := c.Query("accountId"); a != "" {
q = q.Where("account_id = ?", a)
}
var list []models.Challenge
q.Order("id desc").Limit(100).Find(&list)
response.OK(c, gin.H{"list": list})
}
// Solve 人工完成挑战(验证码/APP确认;扫码由 Agent 完成)
func (h *ChallengeHandler) Solve(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req struct {
Value string `json:"value"`
}
_ = c.ShouldBindJSON(&req)
var challenge models.Challenge
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&challenge).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "挑战不存在")
return
}
if challenge.Status != "active" {
response.Fail(c, http.StatusConflict, 3003, "挑战已结束")
return
}
now := time.Now()
if err = h.DB.Model(&challenge).Updates(map[string]interface{}{"status": "solved", "solved_at": &now}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "challenge.solve", "challenge:"+itoa(id), gin.H{"kind": challenge.Kind})
response.OK(c, nil)
}
// Resend 一键重发(重置过期时间,实际重发出 Agent 执行)
func (h *ChallengeHandler) Resend(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var challenge models.Challenge
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&challenge).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "挑战不存在")
return
}
if challenge.Status != "expired" && challenge.Status != "suspended" {
response.Fail(c, http.StatusConflict, 3003, "当前状态不可重发")
return
}
if err = h.DB.Model(&challenge).Updates(map[string]interface{}{
"status": "active", "expires_at": time.Now().Add(30 * time.Minute),
}).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "challenge.resend", "challenge:"+itoa(id), nil)
response.OK(c, nil)
}
// Suspend 挂起(超时/人工)
func (h *ChallengeHandler) Suspend(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var challenge models.Challenge
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&challenge).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "挑战不存在")
return
}
if challenge.Status != "active" {
response.Fail(c, http.StatusConflict, 3003, "挑战已结束")
return
}
if err = h.DB.Model(&challenge).Update("status", "suspended").Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "challenge.suspend", "challenge:"+itoa(id), nil)
response.OK(c, nil)
}
+201
View File
@@ -0,0 +1,201 @@
package handlers
import (
"crypto/sha256"
"encoding/hex"
"io"
"net/http"
"os"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
)
// MaterialHandler 素材处理器(服务器本地盘一期;StorageDriver 二期 OSS)
type MaterialHandler struct {
DB *gorm.DB
Cfg *config.Config
}
// List 素材列表(分页 + 筛选)
func (h *MaterialHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.Material{}).Where("workspace_id = ? and status = ?", c.GetUint64("wsid"), "ready")
if k := c.Query("kind"); k != "" {
q = q.Where("kind = ?", k)
}
if g := c.Query("group"); g != "" {
q = q.Where("`group` = ?", g)
}
var total int64
q.Count(&total)
var list []models.Material
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&list)
response.OK(c, gin.H{"list": list, "total": total, "page": page, "size": size})
}
// Upload multipart 上传:流式计算 sha256,同工作区按 sha256 去重
func (h *MaterialHandler) Upload(c *gin.Context) {
if h.Cfg.MaxUploadBytes > 0 {
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, h.Cfg.MaxUploadBytes)
}
file, header, err := c.Request.FormFile("file")
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "文件过大或缺少文件字段 file")
return
}
defer file.Close()
wsid := c.GetUint64("wsid")
dir := filepath.Join(h.Cfg.StorageDir, itoa(wsid))
if err = os.MkdirAll(dir, 0o755); err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "存储初始化失败")
return
}
ext := strings.ToLower(filepath.Ext(header.Filename))
if ext == "" || len(ext) > 16 {
ext = ".bin"
}
key := filepath.Join(itoa(wsid), uuid.NewString()+ext)
abspath := filepath.Join(h.Cfg.StorageDir, key)
dst, err := os.Create(abspath)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "写盘失败")
return
}
hasher := sha256.New()
size, err := io.Copy(io.MultiWriter(dst, hasher), file)
dst.Close()
if err != nil {
_ = os.Remove(abspath)
response.Fail(c, http.StatusInternalServerError, 5000, "写入失败")
return
}
sum := hex.EncodeToString(hasher.Sum(nil))
// 去重:同工作区同 sha256 直接复用
var existing models.Material
if err = h.DB.Where("workspace_id = ? and sha256 = ? and status = ?", wsid, sum, "ready").First(&existing).Error; err == nil {
_ = os.Remove(abspath)
response.OK(c, gin.H{"material": existing, "dedup": true})
return
}
kind := "image"
if strings.HasPrefix(header.Header.Get("Content-Type"), "video/") {
kind = "video"
}
group := strings.TrimSpace(c.PostForm("group"))
tags := strings.TrimSpace(c.PostForm("tags"))
if tags == "" {
tags = "[]"
}
material := models.Material{
WorkspaceID: wsid,
Name: header.Filename,
Kind: kind,
Size: size,
SHA256: sum,
Mime: header.Header.Get("Content-Type"),
StorageKey: key,
Group: group,
Tags: tags,
Status: "ready",
CreatedBy: c.GetUint64("uid"),
}
if err = h.DB.Create(&material).Error; err != nil {
_ = os.Remove(abspath)
response.Fail(c, http.StatusInternalServerError, 5000, "入库失败")
return
}
response.Audit(c, "material.upload", "material:"+itoa(material.ID), gin.H{"name": header.Filename, "size": size})
response.OK(c, gin.H{"material": material, "dedup": false})
}
// URL 生成签名直链(10 分钟一次性)
func (h *MaterialHandler) URL(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var material models.Material
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&material).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "素材不存在")
return
}
token := models.TransferToken{
MaterialID: material.ID,
Token: uuid.NewString(),
ExpiresAt: time.Now().Add(10 * time.Minute),
}
if err = h.DB.Create(&token).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "生成链接失败")
return
}
response.OK(c, gin.H{"url": h.Cfg.BaseURL + "/api/v1/files/" + token.Token, "expiresAt": token.ExpiresAt})
}
// ServeFile 凭 token 取文件(一次性,10 分钟有效)
func (h *MaterialHandler) ServeFile(c *gin.Context) {
tokenStr := c.Param("token")
var token models.TransferToken
if err := h.DB.Where("token = ?", tokenStr).First(&token).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "链接无效")
return
}
if time.Now().After(token.ExpiresAt) {
h.DB.Delete(&token)
response.Fail(c, http.StatusGone, 3004, "链接已过期")
return
}
var material models.Material
if err := h.DB.First(&material, token.MaterialID).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "素材不存在")
return
}
// 一次性:取用即吊销
_ = h.DB.Delete(&token).Error
abspath := filepath.Join(h.Cfg.StorageDir, material.StorageKey)
clean := filepath.Clean(abspath)
if !strings.HasPrefix(clean, filepath.Clean(h.Cfg.StorageDir)) {
response.Fail(c, http.StatusForbidden, 1003, "非法路径")
return
}
c.FileAttachment(abspath, material.Name)
}
// Delete 删除素材(含文件)
func (h *MaterialHandler) Delete(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var material models.Material
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&material).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "素材不存在")
return
}
abspath := filepath.Join(h.Cfg.StorageDir, material.StorageKey)
_ = os.Remove(abspath)
if err = h.DB.Delete(&material).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "删除失败")
return
}
response.Audit(c, "material.delete", "material:"+itoa(id), nil)
response.OK(c, nil)
}
+190
View File
@@ -0,0 +1,190 @@
package handlers
import (
"net/http"
"strconv"
"strings"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/auth"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
)
// MemberHandler 成员处理器
type MemberHandler struct {
DB *gorm.DB
Cfg *config.Config
}
type inviteReq struct {
Email string `json:"email" binding:"required"`
Role string `json:"role" binding:"required"`
}
type joinReq struct {
Token string `json:"token" binding:"required"`
}
type roleReq struct {
Role string `json:"role" binding:"required"`
}
var memberRoles = map[string]bool{"admin": true, "operator": true, "reviewer": true, "viewer": true}
func isAdminRole(mrole string) bool { return mrole == "owner" || mrole == "admin" }
type memberRow struct {
models.Member
Email string `json:"email"`
Nickname string `json:"nickname"`
}
// List 成员列表(含用户资料)
func (h *MemberHandler) List(c *gin.Context) {
wsid := c.GetUint64("wsid")
var rows []memberRow
h.DB.Table("members").Select("members.*, users.email as email, users.nickname as nickname").
Joins("join users on users.id = members.user_id").
Where("members.workspace_id = ?", wsid).Order("members.id asc").Scan(&rows)
response.OK(c, gin.H{"list": rows})
}
// Invite 生成邀请链接(owner/admin)
func (h *MemberHandler) Invite(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可邀请成员")
return
}
var req inviteReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
req.Role = strings.ToLower(strings.TrimSpace(req.Role))
if !memberRoles[req.Role] {
response.Fail(c, http.StatusBadRequest, 1001, "角色不合法")
return
}
wsid := c.GetUint64("wsid")
var existing memberRow
err := h.DB.Table("members").Select("members.id").
Joins("join users on users.id = members.user_id").
Where("members.workspace_id = ? and users.email = ?", wsid, strings.ToLower(req.Email)).
Scan(&existing).Error
if err == nil && existing.ID > 0 {
response.Fail(c, http.StatusConflict, 2004, "该邮箱已是成员")
return
}
token, err := auth.IssueInvite(h.Cfg.JWTSecret, wsid, strings.ToLower(req.Email), req.Role)
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "生成邀请失败")
return
}
response.Audit(c, "member.invite", "workspace:"+itoa(wsid), gin.H{"email": req.Email, "role": req.Role})
response.OK(c, gin.H{"token": token, "link": h.Cfg.PublicBaseURL + "/invite?t=" + token})
}
// Join 接受邀请(需登录,邮箱须与邀请一致)
func (h *MemberHandler) Join(c *gin.Context) {
var req joinReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
claims, err := auth.ParseInvite(h.Cfg.JWTSecret, req.Token)
if err != nil {
response.Fail(c, http.StatusBadRequest, 2005, "邀请链接无效或已过期")
return
}
var user models.User
if err = h.DB.First(&user, c.GetUint64("uid")).Error; err != nil {
response.Fail(c, http.StatusUnauthorized, 1002, "用户不存在")
return
}
if !strings.EqualFold(user.Email, claims.Email) {
response.Fail(c, http.StatusForbidden, 1003, "邀请链接与当前账号邮箱不匹配")
return
}
member := models.Member{WorkspaceID: claims.WorkspaceID, UserID: user.ID, Role: claims.Role}
err = h.DB.Where("workspace_id = ? and user_id = ?", claims.WorkspaceID, user.ID).First(&models.Member{}).Error
if err == nil {
response.Fail(c, http.StatusConflict, 2004, "你已是该工作空间成员")
return
}
if err = h.DB.Create(&member).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "加入失败")
return
}
response.Audit(c, "member.join", "workspace:"+itoa(claims.WorkspaceID), gin.H{"role": claims.Role})
response.OK(c, nil)
}
// UpdateRole 修改成员角色(owner/admin;不可改 owner)
func (h *MemberHandler) UpdateRole(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可修改角色")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req roleReq
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
req.Role = strings.ToLower(strings.TrimSpace(req.Role))
if !memberRoles[req.Role] && req.Role != "owner" {
response.Fail(c, http.StatusBadRequest, 1001, "角色不合法")
return
}
var member models.Member
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&member).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "成员不存在")
return
}
if member.Role == "owner" {
response.Fail(c, http.StatusForbidden, 1003, "不可修改 owner 角色")
return
}
if err = h.DB.Model(&member).Update("role", req.Role).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "member.role", "member:"+itoa(id), gin.H{"role": req.Role})
response.OK(c, nil)
}
// Remove 移除成员(owner/admin;不可移除 owner)
func (h *MemberHandler) Remove(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可移除成员")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var member models.Member
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&member).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "成员不存在")
return
}
if member.Role == "owner" {
response.Fail(c, http.StatusForbidden, 1003, "不可移除 owner")
return
}
if err = h.DB.Delete(&member).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "移除失败")
return
}
response.Audit(c, "member.remove", "member:"+itoa(id), nil)
response.OK(c, nil)
}
@@ -0,0 +1,57 @@
package handlers
import (
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// NotificationHandler 站内通知处理器
type NotificationHandler struct {
DB *gorm.DB
}
// List 我的通知 + 未读数
func (h *NotificationHandler) List(c *gin.Context) {
q := h.DB.Model(&models.Notification{}).Where("workspace_id = ? and user_id = ?", c.GetUint64("wsid"), c.GetUint64("uid"))
if k := c.Query("kind"); k != "" {
q = q.Where("kind = ?", k)
}
var unread int64
q.Where("`read` = ?", false).Count(&unread)
var list []models.Notification
q.Order("id desc").Limit(100).Find(&list)
response.OK(c, gin.H{"list": list, "unread": unread})
}
// Read 标记已读
func (h *NotificationHandler) Read(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
if err = h.DB.Model(&models.Notification{}).
Where("id = ? and user_id = ?", id, c.GetUint64("uid")).
Update("read", true).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.OK(c, nil)
}
// ReadAll 全部已读
func (h *NotificationHandler) ReadAll(c *gin.Context) {
if err := h.DB.Model(&models.Notification{}).
Where("workspace_id = ? and user_id = ?", c.GetUint64("wsid"), c.GetUint64("uid")).
Update("read", true).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.OK(c, nil)
}
+25
View File
@@ -0,0 +1,25 @@
package handlers
import (
"gorm.io/gorm"
"everypublish/server/internal/models"
)
// Notify 写入站内通知(批量用户)
func Notify(db *gorm.DB, wsid uint64, userIDs []uint64, kind, title, content string) error {
if len(userIDs) == 0 {
return nil
}
rows := make([]models.Notification, 0, len(userIDs))
for _, uid := range userIDs {
rows = append(rows, models.Notification{
WorkspaceID: wsid,
UserID: uid,
Kind: kind,
Title: title,
Content: content,
})
}
return db.Create(&rows).Error
}
+279
View File
@@ -0,0 +1,279 @@
package handlers
import (
"encoding/json"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
"everypublish/server/internal/task"
"everypublish/server/internal/ws"
)
// TaskHandler 发布任务处理器(状态机见 internal/task)
type TaskHandler struct {
DB *gorm.DB
Dispatcher *ws.Dispatcher
}
type taskReq struct {
Title string `json:"title" binding:"required"`
Content string `json:"content"`
Tags []string `json:"tags"`
AccountIDs []uint64 `json:"accountIds"`
MaterialIDs []uint64 `json:"materialIds"`
ScheduleAt *int64 `json:"scheduleAt"`
Priority int `json:"priority"`
}
func marshalJSONArray[T any](arr []T) string {
raw, _ := json.Marshal(arr)
if raw == nil {
return "[]"
}
return string(raw)
}
// List 任务列表(状态筛选 + 分页)
func (h *TaskHandler) List(c *gin.Context) {
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
if page < 1 {
page = 1
}
if size < 1 || size > 100 {
size = 20
}
q := h.DB.Model(&models.Task{}).Where("workspace_id = ?", c.GetUint64("wsid"))
if s := c.Query("status"); s != "" {
q = q.Where("status = ?", s)
}
if s := c.Query("schedule"); s != "" {
if s == "scheduled" {
q = q.Where("schedule_at is not null and schedule_at > ?", time.Now())
} else if s == "immediate" {
q = q.Where("schedule_at is null or schedule_at <= ?", time.Now())
}
}
var total int64
q.Count(&total)
var list []models.Task
q.Order("id desc").Offset((page - 1) * size).Limit(size).Find(&list)
response.OK(c, gin.H{"list": list, "total": total, "page": page, "size": size})
}
// Create 新建任务(draft)
func (h *TaskHandler) Create(c *gin.Context) {
var req taskReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法:"+err.Error())
return
}
if len(req.AccountIDs) == 0 {
response.Fail(c, http.StatusBadRequest, 1001, "至少选择一个平台账号")
return
}
if len(req.MaterialIDs) == 0 {
response.Fail(c, http.StatusBadRequest, 1001, "至少选择一个素材")
return
}
priority := req.Priority
if priority < 1 || priority > 10 {
priority = 5
}
row := models.Task{
WorkspaceID: c.GetUint64("wsid"),
Title: req.Title,
Content: req.Content,
Tags: marshalJSONArray(req.Tags),
AccountIDs: marshalJSONArray(req.AccountIDs),
MaterialIDs: marshalJSONArray(req.MaterialIDs),
Status: task.Draft,
Priority: priority,
MaxRetry: 3,
CreatedBy: c.GetUint64("uid"),
}
if req.ScheduleAt != nil {
ts := time.UnixMilli(*req.ScheduleAt)
row.ScheduleAt = &ts
}
if err := h.DB.Create(&row).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "创建失败")
return
}
response.Audit(c, "task.create", "task:"+itoa(row.ID), gin.H{"title": req.Title})
response.OK(c, row)
}
// Get 任务详情
func (h *TaskHandler) Get(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var row models.Task
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&row).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "任务不存在")
return
}
response.OK(c, row)
}
// Update 编辑草稿(仅 draft/rejected)
func (h *TaskHandler) Update(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req taskReq
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var row models.Task
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&row).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "任务不存在")
return
}
if row.Status != task.Draft && row.Status != task.Rejected {
response.Fail(c, http.StatusConflict, 3005, "仅草稿或已驳回任务可编辑")
return
}
updates := map[string]interface{}{
"title": req.Title,
"content": req.Content,
"tags": marshalJSONArray(req.Tags),
"account_ids": marshalJSONArray(req.AccountIDs),
"material_ids": marshalJSONArray(req.MaterialIDs),
}
if req.ScheduleAt != nil {
ts := time.UnixMilli(*req.ScheduleAt)
updates["schedule_at"] = &ts
}
if err = h.DB.Model(&row).Updates(updates).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "task.update", "task:"+itoa(id), gin.H{"title": req.Title})
response.OK(c, nil)
}
// Delete 删除草稿
func (h *TaskHandler) Delete(c *gin.Context) {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var row models.Task
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&row).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "任务不存在")
return
}
if row.Status != task.Draft {
response.Fail(c, http.StatusConflict, 3005, "仅草稿可删除")
return
}
if err = h.DB.Delete(&row).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "删除失败")
return
}
response.Audit(c, "task.delete", "task:"+itoa(id), nil)
response.OK(c, nil)
}
// apply 状态转移并落库
func (h *TaskHandler) apply(c *gin.Context, action task.Action, extra map[string]interface{}) *models.Task {
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return nil
}
var row models.Task
if err = h.DB.Where("id = ? and workspace_id = ?", id, c.GetUint64("wsid")).First(&row).Error; err != nil {
response.Fail(c, http.StatusNotFound, 1004, "任务不存在")
return nil
}
next, ok := task.Next(row.Status, action)
if !ok {
response.Fail(c, http.StatusConflict, 3005, "当前状态("+row.Status+")不允许该操作("+string(action)+")")
return nil
}
updates := map[string]interface{}{"status": next}
for k, v := range extra {
updates[k] = v
}
if err = h.DB.Model(&row).Updates(updates).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return nil
}
row.Status = next
return &row
}
// Submit 提交审核
func (h *TaskHandler) Submit(c *gin.Context) {
if row := h.apply(c, task.ActionSubmit, nil); row != nil {
response.Audit(c, "task.submit", "task:"+itoa(row.ID), nil)
response.OK(c, row)
}
}
// Approve 审核通过 → queued
func (h *TaskHandler) Approve(c *gin.Context) {
if row := h.apply(c, task.ActionApprove, nil); row != nil {
_ = Notify(h.DB, c.GetUint64("wsid"), []uint64{row.CreatedBy}, "task", "任务已通过审核", row.Title)
response.Audit(c, "task.approve", "task:"+itoa(row.ID), nil)
if h.Dispatcher != nil {
_ = h.Dispatcher.TryDispatch(row.ID) // 在线设备立即下发
}
response.OK(c, row)
}
}
// Reject 驳回(附批注)
func (h *TaskHandler) Reject(c *gin.Context) {
var req struct {
Note string `json:"note"`
}
_ = c.ShouldBindJSON(&req)
if row := h.apply(c, task.ActionReject, map[string]interface{}{"review_note": req.Note, "reviewed_by": c.GetUint64("uid")}); row != nil {
_ = Notify(h.DB, c.GetUint64("wsid"), []uint64{row.CreatedBy}, "task", "任务被驳回", req.Note)
response.Audit(c, "task.reject", "task:"+itoa(row.ID), gin.H{"note": req.Note})
response.OK(c, row)
}
}
// Resubmit 驳回后重新提交
func (h *TaskHandler) Resubmit(c *gin.Context) {
if row := h.apply(c, task.ActionResubmit, nil); row != nil {
response.Audit(c, "task.resubmit", "task:"+itoa(row.ID), nil)
response.OK(c, row)
}
}
// Retry 失败重试
func (h *TaskHandler) Retry(c *gin.Context) {
if row := h.apply(c, task.ActionRetry, map[string]interface{}{"error_message": ""}); row != nil {
response.Audit(c, "task.retry", "task:"+itoa(row.ID), nil)
if h.Dispatcher != nil {
_ = h.Dispatcher.TryDispatch(row.ID)
}
response.OK(c, row)
}
}
// Cancel 取消
func (h *TaskHandler) Cancel(c *gin.Context) {
if row := h.apply(c, task.ActionCancel, nil); row != nil {
response.Audit(c, "task.cancel", "task:"+itoa(row.ID), nil)
response.OK(c, row)
}
}
+7
View File
@@ -0,0 +1,7 @@
package handlers
import "strconv"
func itoa(v uint64) string {
return strconv.FormatUint(v, 10)
}
+84
View File
@@ -0,0 +1,84 @@
package handlers
import (
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// WorkspaceHandler 工作空间处理器
type WorkspaceHandler struct {
DB *gorm.DB
}
type workspaceReq struct {
Name string `json:"name" binding:"required"`
}
// Create 创建工作空间(创建者为 owner)
func (h *WorkspaceHandler) Create(c *gin.Context) {
var req workspaceReq
if err := c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var ws models.Workspace
err := h.DB.Transaction(func(tx *gorm.DB) error {
ws = models.Workspace{Name: req.Name, OwnerID: c.GetUint64("uid"), Plan: "free", Status: "active"}
if err := tx.Create(&ws).Error; err != nil {
return err
}
member := models.Member{WorkspaceID: ws.ID, UserID: c.GetUint64("uid"), Role: "owner"}
return tx.Create(&member).Error
})
if err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "创建失败")
return
}
response.Audit(c, "workspace.create", "workspace:"+itoa(ws.ID), gin.H{"name": req.Name})
response.OK(c, ws)
}
// List 我所在的工作空间
func (h *WorkspaceHandler) List(c *gin.Context) {
var members []models.Member
h.DB.Where("user_id = ?", c.GetUint64("uid")).Find(&members)
ids := make([]uint64, 0, len(members))
for _, m := range members {
ids = append(ids, m.WorkspaceID)
}
var list []models.Workspace
if len(ids) > 0 {
h.DB.Where("id in ?", ids).Order("id asc").Find(&list)
}
response.OK(c, gin.H{"list": list})
}
// Update 修改工作空间资料(owner/admin)
func (h *WorkspaceHandler) Update(c *gin.Context) {
if !isAdminRole(c.GetString("mrole")) {
response.Fail(c, http.StatusForbidden, 1003, "仅管理员可修改工作空间")
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
var req workspaceReq
if err = c.ShouldBindJSON(&req); err != nil {
response.Fail(c, http.StatusBadRequest, 1001, "参数不合法")
return
}
if err = h.DB.Model(&models.Workspace{}).Where("id = ?", id).Update("name", req.Name).Error; err != nil {
response.Fail(c, http.StatusInternalServerError, 5000, "更新失败")
return
}
response.Audit(c, "workspace.update", "workspace:"+itoa(id), gin.H{"name": req.Name})
response.OK(c, nil)
}
+38
View File
@@ -0,0 +1,38 @@
package middleware
import (
"encoding/json"
"net/http"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/response"
"everypublish/server/internal/models"
)
// Audit 审计埋点:处理器调用 response.Audit(c, action, resource, detail) 声明,
// 本中间件在请求成功(2xx)后落库。
func Audit(db *gorm.DB) gin.HandlerFunc {
return func(c *gin.Context) {
c.Next()
metaRaw, exists := c.Get("audit")
if !exists {
return
}
meta, ok := metaRaw.(response.AuditMeta)
if !ok || c.Writer.Status() >= http.StatusBadRequest {
return
}
detail, _ := json.Marshal(meta.Detail)
log := models.AuditLog{
WorkspaceID: c.GetUint64("wsid"),
UserID: c.GetUint64("uid"),
Action: meta.Action,
Resource: meta.Resource,
Detail: string(detail),
IP: c.ClientIP(),
}
_ = db.Create(&log).Error
}
}
+47
View File
@@ -0,0 +1,47 @@
package middleware
import (
"net/http"
"strings"
"github.com/gin-gonic/gin"
"everypublish/server/internal/api/response"
"everypublish/server/internal/auth"
"everypublish/server/internal/config"
)
// AuthRequired 校验访问令牌并把声明写入上下文
func AuthRequired(cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
header := c.GetHeader("Authorization")
if !strings.HasPrefix(header, "Bearer ") {
response.Fail(c, http.StatusUnauthorized, 1002, "未登录或令牌缺失")
c.Abort()
return
}
claims, err := auth.ParseAccess(cfg.JWTSecret, strings.TrimPrefix(header, "Bearer "))
if err != nil {
response.Fail(c, http.StatusUnauthorized, 1002, "令牌无效或已过期")
c.Abort()
return
}
c.Set("uid", claims.UID)
c.Set("wsid", claims.WorkspaceID)
c.Set("role", claims.Role)
c.Set("mrole", claims.MemberRole)
c.Next()
}
}
// AdminRequired 平台管理员守卫(D7 挂载 admin 路由)
func AdminRequired() gin.HandlerFunc {
return func(c *gin.Context) {
if c.GetString("role") != "admin" {
response.Fail(c, http.StatusForbidden, 1003, "无平台管理权限")
c.Abort()
return
}
c.Next()
}
}
+29
View File
@@ -0,0 +1,29 @@
package response
import (
"net/http"
"github.com/gin-gonic/gin"
)
// OK 统一成功返回
func OK(c *gin.Context, data interface{}) {
c.JSON(http.StatusOK, gin.H{"code": 0, "message": "ok", "data": data})
}
// Fail 统一失败返回
func Fail(c *gin.Context, httpCode int, bizCode int, msg string) {
c.JSON(httpCode, gin.H{"code": bizCode, "message": msg, "data": nil})
}
// AuditMeta 处理器埋点元数据
type AuditMeta struct {
Action string `json:"action"`
Resource string `json:"resource"`
Detail interface{} `json:"detail"`
}
// Audit 审计埋点声明(由 middleware.Audit 中间件落库)
func Audit(c *gin.Context, action, resource string, detail interface{}) {
c.Set("audit", AuditMeta{Action: action, Resource: resource, Detail: detail})
}
+156
View File
@@ -0,0 +1,156 @@
package api
import (
"crypto/ed25519"
"net/http"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/api/handlers"
"everypublish/server/internal/api/middleware"
"everypublish/server/internal/cache"
"everypublish/server/internal/config"
"everypublish/server/internal/ws"
)
// Router 装配 gin 路由
func Router(cfg *config.Config, db *gorm.DB, rds *cache.Redis, hub *ws.Hub, serverKey ed25519.PrivateKey, dsp *ws.Dispatcher) *gin.Engine {
r := gin.New()
r.Use(gin.Logger(), gin.Recovery())
r.GET("/health", func(c *gin.Context) {
dbOK := false
if sqlDB, err := db.DB(); err == nil {
dbOK = sqlDB.Ping() == nil
}
redisOK := rds.Ping() == nil
c.JSON(http.StatusOK, gin.H{
"status": "ok",
"service": "everypublish-server",
"time": time.Now().Format(time.RFC3339),
"db": dbOK,
"redis": redisOK,
})
})
// WSS 双通道:Agent 设备签名 / 浏览器 JWT
r.GET("/ws/agent", ws.HandleAgent(cfg, db, hub, serverKey))
r.GET("/ws/browser", ws.HandleBrowser(cfg, hub))
api := r.Group("/api/v1")
api.Use(middleware.Audit(db))
authH := &handlers.AuthHandler{DB: db, Rds: rds, Cfg: cfg}
wsH := &handlers.WorkspaceHandler{DB: db}
mbH := &handlers.MemberHandler{DB: db, Cfg: cfg}
adH := &handlers.AuditHandler{DB: db}
accH := &handlers.AccountHandler{DB: db, Hub: hub}
matH := &handlers.MaterialHandler{DB: db, Cfg: cfg}
tskH := &handlers.TaskHandler{DB: db, Dispatcher: dsp}
chlH := &handlers.ChallengeHandler{DB: db}
ntfH := &handlers.NotificationHandler{DB: db}
agH := &handlers.AgentHandler{DB: db, Hub: hub}
auth := api.Group("/auth")
{
auth.POST("/register", authH.Register)
auth.POST("/login", authH.Login)
auth.POST("/refresh", authH.Refresh)
auth.POST("/logout", middleware.AuthRequired(cfg), authH.Logout)
auth.GET("/me", middleware.AuthRequired(cfg), authH.Me)
}
wsg := api.Group("/workspaces", middleware.AuthRequired(cfg))
{
wsg.POST("", wsH.Create)
wsg.GET("", wsH.List)
wsg.PUT("/:id", wsH.Update)
}
mb := api.Group("/members", middleware.AuthRequired(cfg))
{
mb.GET("", mbH.List)
mb.POST("/invite", mbH.Invite)
mb.POST("/join", mbH.Join)
mb.PUT("/:id/role", mbH.UpdateRole)
mb.DELETE("/:id", mbH.Remove)
}
ad := api.Group("/audit-logs", middleware.AuthRequired(cfg))
{
ad.GET("", adH.List)
}
acc := api.Group("/accounts", middleware.AuthRequired(cfg))
{
acc.GET("", accH.List)
acc.POST("", accH.Create)
acc.PUT("/:id", accH.Update)
acc.DELETE("/:id", accH.Delete)
acc.POST("/:id/bind", accH.Bind)
}
mat := api.Group("/materials", middleware.AuthRequired(cfg))
{
mat.GET("", matH.List)
mat.POST("", matH.Upload)
mat.DELETE("/:id", matH.Delete)
mat.GET("/:id/url", matH.URL)
}
files := api.Group("/files")
{
files.GET("/:token", matH.ServeFile)
}
tsk := api.Group("/tasks", middleware.AuthRequired(cfg))
{
tsk.GET("", tskH.List)
tsk.POST("", tskH.Create)
tsk.GET("/:id", tskH.Get)
tsk.PUT("/:id", tskH.Update)
tsk.DELETE("/:id", tskH.Delete)
tsk.POST("/:id/submit", tskH.Submit)
tsk.POST("/:id/approve", tskH.Approve)
tsk.POST("/:id/reject", tskH.Reject)
tsk.POST("/:id/resubmit", tskH.Resubmit)
tsk.POST("/:id/retry", tskH.Retry)
tsk.POST("/:id/cancel", tskH.Cancel)
}
chl := api.Group("/challenges", middleware.AuthRequired(cfg))
{
chl.GET("", chlH.List)
chl.POST("/:id/solve", chlH.Solve)
chl.POST("/:id/resend", chlH.Resend)
chl.POST("/:id/suspend", chlH.Suspend)
}
ntf := api.Group("/notifications", middleware.AuthRequired(cfg))
{
ntf.GET("", ntfH.List)
ntf.POST("/:id/read", ntfH.Read)
ntf.POST("/read-all", ntfH.ReadAll)
}
ag := api.Group("/agent")
{
ag.POST("/pair-code", middleware.AuthRequired(cfg), agH.PairCode)
ag.GET("/devices", middleware.AuthRequired(cfg), agH.Devices)
ag.PUT("/devices/:id/revoke", middleware.AuthRequired(cfg), agH.Revoke)
}
api.POST("/agent/pair", agH.Pair)
admH := &handlers.AdminHandler{DB: db, Hub: hub}
adm := api.Group("/admin", middleware.AuthRequired(cfg), middleware.AdminRequired())
{
adm.GET("/tenants", admH.Tenants)
adm.PUT("/tenants/:id/status", admH.TenantStatus)
adm.GET("/agents", admH.Agents)
adm.PUT("/agents/:id/revoke", admH.AgentRevoke)
}
return r
}
+250
View File
@@ -0,0 +1,250 @@
package api_test
import (
"bytes"
"crypto/ed25519"
"crypto/rand"
"encoding/json"
"fmt"
"net/http/httptest"
"os"
"testing"
"github.com/gin-gonic/gin"
gormmysql "gorm.io/driver/mysql"
"gorm.io/gorm"
"everypublish/server/internal/api"
"everypublish/server/internal/cache"
"everypublish/server/internal/config"
"everypublish/server/internal/db"
"everypublish/server/internal/models"
"everypublish/server/internal/ws"
)
var (
testRouter *gin.Engine
testDB *gorm.DB
testHub *ws.Hub
)
func TestMain(m *testing.M) {
adminDSN := os.Getenv("MYSQL_ADMIN_DSN")
if adminDSN == "" {
adminDSN = "root:everypublish@tcp(127.0.0.1:3306)/?charset=utf8mb4&parseTime=True&loc=Local"
}
adm, err := gorm.Open(gormmysql.Open(adminDSN), &gorm.Config{})
if err != nil {
fmt.Println("skip: mysql admin connect failed:", err)
os.Exit(0)
}
if err = adm.Exec("CREATE DATABASE IF NOT EXISTS everypublish_test CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci").Error; err != nil {
fmt.Println("skip: create test db failed:", err)
os.Exit(0)
}
testDSN := "root:everypublish@tcp(127.0.0.1:3306)/everypublish_test?charset=utf8mb4&parseTime=True&loc=Local"
testDB, err = gorm.Open(gormmysql.Open(testDSN), &gorm.Config{})
if err != nil {
fmt.Println("skip: test db connect failed:", err)
os.Exit(0)
}
resetTestDB()
cfg := &config.Config{
ServerAddr: ":0",
MySQLDSN: testDSN,
RedisAddr: "127.0.0.1:6379",
JWTSecret: "test-secret",
BaseURL: "http://127.0.0.1:8090",
StorageDir: "/tmp/everypublish-test",
}
rds := cache.New(cfg)
testHub = ws.NewHub()
dsp := ws.NewDispatcher(testDB, testHub, cfg)
_, serverKey, _ := ed25519.GenerateKey(rand.Reader)
testRouter = api.Router(cfg, testDB, rds, testHub, serverKey, dsp)
code := m.Run()
os.Exit(code)
}
// resetTestDB 清空并重建表 + 清空连接中枢(保证 -count 多轮可重复)
func resetTestDB() {
if testHub != nil {
testHub.Clear()
}
for _, t := range []interface{}{
&models.TransferToken{}, &models.PairingCode{}, &models.Notification{}, &models.AuditLog{},
&models.Challenge{}, &models.AgentDevice{}, &models.Task{}, &models.Material{},
&models.Account{}, &models.Member{}, &models.Workspace{}, &models.User{},
} {
testDB.Migrator().DropTable(t)
}
_ = db.AutoMigrate(testDB)
}
type envelope struct {
Code int `json:"code"`
Message string `json:"message"`
Data json.RawMessage `json:"data"`
}
func doReq(t *testing.T, method, path, token string, body interface{}) (int, envelope) {
t.Helper()
var buf bytes.Buffer
if body != nil {
raw, _ := json.Marshal(body)
buf.Write(raw)
}
req := httptest.NewRequest(method, path, &buf)
req.Header.Set("Content-Type", "application/json")
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
w := httptest.NewRecorder()
testRouter.ServeHTTP(w, req)
var env envelope
_ = json.Unmarshal(w.Body.Bytes(), &env)
return w.Code, env
}
type tokens struct {
AccessToken string `json:"accessToken"`
RefreshToken string `json:"refreshToken"`
}
func register(t *testing.T, email string) (string, string) {
t.Helper()
code, env := doReq(t, "POST", "/api/v1/auth/register", "", gin.H{"email": email, "password": "password-123", "nickname": "测试用户"})
if code != 200 || env.Code != 0 {
t.Fatalf("register failed: %d %s", code, env.Message)
}
var tk tokens
_ = json.Unmarshal(env.Data, &tk)
return tk.AccessToken, tk.RefreshToken
}
func TestAuthFlow(t *testing.T) {
resetTestDB()
access, refresh := register(t, "alice@test.com")
// 登录
code, env := doReq(t, "POST", "/api/v1/auth/login", "", gin.H{"email": "alice@test.com", "password": "password-123"})
if code != 200 || env.Code != 0 {
t.Fatalf("login failed: %d %s", code, env.Message)
}
// me
code, env = doReq(t, "GET", "/api/v1/auth/me", access, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("me failed: %d %s", code, env.Message)
}
// 错误密码
code, env = doReq(t, "POST", "/api/v1/auth/login", "", gin.H{"email": "alice@test.com", "password": "wrong-pass-1"})
if code != 401 || env.Code != 2002 {
t.Fatalf("wrong password should 2002, got %d %d", code, env.Code)
}
// 刷新轮换
code, env = doReq(t, "POST", "/api/v1/auth/refresh", "", gin.H{"refreshToken": refresh})
if code != 200 || env.Code != 0 {
t.Fatalf("refresh failed: %d %s", code, env.Message)
}
var tk tokens
_ = json.Unmarshal(env.Data, &tk)
newAccess, newRefresh := tk.AccessToken, tk.RefreshToken
// 旧 refresh 已被吊销
code, env = doReq(t, "POST", "/api/v1/auth/refresh", "", gin.H{"refreshToken": refresh})
if code != 401 || env.Code != 2003 {
t.Fatalf("old refresh should be revoked, got %d %d", code, env.Code)
}
// 登出后 refresh 失效
code, _ = doReq(t, "POST", "/api/v1/auth/logout", newAccess, gin.H{"refreshToken": newRefresh})
if code != 200 {
t.Fatalf("logout failed: %d", code)
}
code, env = doReq(t, "POST", "/api/v1/auth/refresh", "", gin.H{"refreshToken": newRefresh})
if code != 401 || env.Code != 2003 {
t.Fatalf("refresh after logout should fail, got %d %d", code, env.Code)
}
}
func TestWorkspaceAndMemberFlow(t *testing.T) {
resetTestDB()
aAccess, _ := register(t, "boss@test.com")
bAccess, _ := register(t, "worker@test.com")
// A 建第二个工作空间
code, env := doReq(t, "POST", "/api/v1/workspaces", aAccess, gin.H{"name": "第二空间"})
if code != 200 || env.Code != 0 {
t.Fatalf("create ws failed: %d %s", code, env.Message)
}
// A 邀请 B(viewer)
code, env = doReq(t, "POST", "/api/v1/members/invite", aAccess, gin.H{"email": "worker@test.com", "role": "viewer"})
if code != 200 || env.Code != 0 {
t.Fatalf("invite failed: %d %s", code, env.Message)
}
var inv struct {
Token string `json:"token"`
Link string `json:"link"`
}
_ = json.Unmarshal(env.Data, &inv)
if inv.Token == "" || inv.Link == "" {
t.Fatal("invite token/link missing")
}
// B 用错误邮箱邀请应失败(换 A 邀请 C 不存在的邮箱,B 接受会失败)
code, env = doReq(t, "POST", "/api/v1/members/join", bAccess, gin.H{"token": inv.Token})
if code != 200 || env.Code != 0 {
t.Fatalf("join failed: %d %s", code, env.Message)
}
// A 查成员:2 人
code, env = doReq(t, "GET", "/api/v1/members", aAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("list members failed: %d %s", code, env.Message)
}
var ml struct {
List []models.Member `json:"list"`
}
_ = json.Unmarshal(env.Data, &ml)
if len(ml.List) != 2 {
t.Fatalf("expect 2 members, got %d", len(ml.List))
}
var bMember uint64
var ownerMember uint64
for _, m := range ml.List {
if m.Role == "viewer" {
bMember = m.ID
}
if m.Role == "owner" {
ownerMember = m.ID
}
}
// A 改 B 为 operator
code, env = doReq(t, "PUT", fmt.Sprintf("/api/v1/members/%d/role", bMember), aAccess, gin.H{"role": "operator"})
if code != 200 || env.Code != 0 {
t.Fatalf("update role failed: %d %s", code, env.Message)
}
// owner 不可被移除
code, _ = doReq(t, "DELETE", fmt.Sprintf("/api/v1/members/%d", ownerMember), aAccess, nil)
if code != 403 {
t.Fatalf("remove owner should 403, got %d", code)
}
// owner 角色不可被修改
code, _ = doReq(t, "PUT", fmt.Sprintf("/api/v1/members/%d/role", ownerMember), aAccess, gin.H{"role": "viewer"})
if code != 403 {
t.Fatalf("change owner role should 403, got %d", code)
}
// A 移除 B
code, env = doReq(t, "DELETE", fmt.Sprintf("/api/v1/members/%d", bMember), aAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("remove member failed: %d %s", code, env.Message)
}
// 审计日志存在(登录/建空间/邀请/改角色/移除等)
code, env = doReq(t, "GET", "/api/v1/audit-logs?size=50", aAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("audit list failed: %d %s", code, env.Message)
}
var al struct {
List []models.AuditLog `json:"list"`
Total int64 `json:"total"`
}
_ = json.Unmarshal(env.Data, &al)
if al.Total < 3 {
t.Fatalf("expect >=3 audit logs, got %d", al.Total)
}
}
+247
View File
@@ -0,0 +1,247 @@
package api_test
import (
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/json"
"fmt"
"net/http/httptest"
"strconv"
"strings"
"testing"
"time"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"everypublish/server/internal/models"
"everypublish/shared/proto"
)
func readEnv(t *testing.T, conn *websocket.Conn) *proto.Envelope {
t.Helper()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_, raw, err := conn.Read(ctx)
if err != nil {
t.Fatalf("ws read: %v", err)
}
var env proto.Envelope
if err = json.Unmarshal(raw, &env); err != nil {
t.Fatalf("ws unmarshal: %v", err)
}
return &env
}
func writeEnv(t *testing.T, conn *websocket.Conn, env *proto.Envelope) {
t.Helper()
raw, _ := json.Marshal(env)
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
if err := conn.Write(ctx, websocket.MessageText, raw); err != nil {
t.Fatalf("ws write: %v", err)
}
}
func TestWSAgentFullLoop(t *testing.T) {
resetTestDB()
access, _ := register(t, "agentuser@test.com")
// 1. 配对
code, env := doReq(t, "POST", "/api/v1/agent/pair-code", access, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("pair-code failed: %d %s", code, env.Message)
}
var pc struct {
Code string `json:"code"`
}
_ = json.Unmarshal(env.Data, &pc)
pub, priv, _ := ed25519.GenerateKey(rand.Reader)
pubB64 := base64.StdEncoding.EncodeToString(pub)
code, env = doReq(t, "POST", "/api/v1/agent/pair", "", gin.H{
"code": pc.Code, "deviceName": "test-device", "os": "test", "version": "0.0.1", "publicKey": pubB64,
})
if code != 200 || env.Code != 0 {
t.Fatalf("pair failed: %d %s", code, env.Message)
}
var pd struct {
DeviceID uint64 `json:"deviceId"`
}
_ = json.Unmarshal(env.Data, &pd)
if pd.DeviceID == 0 {
t.Fatal("deviceId missing")
}
// 2. 坏签名 hello 应被拒
srv := httptest.NewServer(testRouter)
defer srv.Close()
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws/agent"
badConn, _, err := websocket.Dial(context.Background(), wsURL, nil)
if err != nil {
t.Fatalf("dial: %v", err)
}
badHello := proto.NewEnvelope("id1", proto.TypeHello, proto.DeviceHello{
DeviceID: strconv.FormatUint(pd.DeviceID, 10),
Nonce: "n1",
TS: time.Now().UnixMilli(),
Sig: base64.StdEncoding.EncodeToString([]byte("bad-sig")),
Version: "0.0.1",
})
rawBad, _ := json.Marshal(badHello)
_ = badConn.Write(context.Background(), websocket.MessageText, rawBad)
ctx2, cancel2 := context.WithTimeout(context.Background(), 3*time.Second)
_, _, err = badConn.Read(ctx2)
cancel2()
if err == nil {
t.Fatal("bad signature should be rejected")
}
// 3. 正常握手
conn, _, err := websocket.Dial(context.Background(), wsURL, nil)
if err != nil {
t.Fatalf("dial: %v", err)
}
nonce := "agent-test-nonce"
ts := time.Now().UnixMilli()
sig := ed25519.Sign(priv, []byte(nonce+"|"+strconv.FormatInt(ts, 10)))
writeEnv(t, conn, proto.NewEnvelope("id2", proto.TypeHello, proto.DeviceHello{
DeviceID: strconv.FormatUint(pd.DeviceID, 10),
Nonce: nonce,
TS: ts,
Sig: base64.StdEncoding.EncodeToString(sig),
Version: "0.0.1",
}))
ackEnv := readEnv(t, conn)
if ackEnv.Type != proto.TypeHelloAck {
t.Fatalf("expect hello.ack, got %s", ackEnv.Type)
}
var ack proto.HelloAck
_ = json.Unmarshal(ackEnv.Payload, &ack)
if ack.SessionID == "" || ack.ServerNonce == "" {
t.Fatal("hello.ack missing fields")
}
// 4. 心跳 → heartbeat.ack
writeEnv(t, conn, proto.NewEnvelope("id3", proto.TypeHeartbeat, proto.Heartbeat{DeviceID: "1", TS: time.Now().UnixMilli()}))
hbAck := readEnv(t, conn)
if hbAck.Type != proto.TypeHeartbeatAck {
t.Fatalf("expect heartbeat.ack, got %s", hbAck.Type)
}
// 5. 账号绑定 → challenge.new → solve → active
code, env = doReq(t, "POST", "/api/v1/accounts", access, gin.H{"platform": "douyin", "accountName": "抖音号01"})
if code != 200 {
t.Fatalf("create account failed: %d", code)
}
var acc models.Account
_ = json.Unmarshal(env.Data, &acc)
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/accounts/%d/bind", acc.ID), access, nil)
if code != 200 {
t.Fatalf("bind failed: %d", code)
}
chEnv := readEnv(t, conn)
if chEnv.Type != proto.TypeChallenge {
t.Fatalf("expect challenge.new, got %s", chEnv.Type)
}
var ch proto.Challenge
_ = json.Unmarshal(chEnv.Payload, &ch)
if ch.AccountID == "" {
t.Fatal("challenge missing accountId")
}
writeEnv(t, conn, proto.NewEnvelope("id5", proto.TypeChallengeSolve, proto.ChallengeSolve{ChallengeID: ch.ChallengeID, Value: "fake-token"}))
time.Sleep(200 * time.Millisecond)
code, env = doReq(t, "GET", "/api/v1/accounts?status=active", access, nil)
if code != 200 {
t.Fatalf("accounts list failed: %d", code)
}
var al struct {
List []models.Account `json:"list"`
}
_ = json.Unmarshal(env.Data, &al)
if len(al.List) != 1 || al.List[0].AgentDeviceID != pd.DeviceID {
t.Fatalf("account should be active bound to device, got %+v", al.List)
}
// 6. 任务闭环(含下发延迟测量)+ 浏览器实时事件
browserConn, _, err := websocket.Dial(context.Background(), "ws"+strings.TrimPrefix(srv.URL, "http")+"/ws/browser?token="+access, nil)
if err != nil {
t.Fatalf("browser dial: %v", err)
}
matID, _ := uploadMaterial(t, access, "loop.mp4", []byte("loop-content"))
code, env = doReq(t, "POST", "/api/v1/tasks", access, gin.H{
"title": "实时闭环任务", "content": "正文",
"accountIds": []uint64{acc.ID}, "materialIds": []uint64{matID},
})
if code != 200 {
t.Fatalf("create task failed: %d", code)
}
var taskRow models.Task
_ = json.Unmarshal(env.Data, &taskRow)
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/submit", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("submit failed: %d", code)
}
t0 := time.Now()
code, _ = doReq(t, "POST", fmt.Sprintf("/api/v1/tasks/%d/approve", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("approve failed: %d", code)
}
pushEnv := readEnv(t, conn)
latency := time.Since(t0)
t.Logf("下发延迟 approve→task.push = %s", latency)
if latency > 2*time.Second {
t.Fatalf("dispatch latency too high: %s", latency)
}
if pushEnv.Type != proto.TypeTaskPush {
t.Fatalf("expect task.push, got %s", pushEnv.Type)
}
var push proto.TaskPush
_ = json.Unmarshal(pushEnv.Payload, &push)
if push.TaskID == "" || len(push.MaterialURLs) == 0 {
t.Fatalf("task.push missing fields: %+v", push)
}
writeEnv(t, conn, proto.NewEnvelope(pushEnv.ID, proto.TypeTaskAck, proto.TaskAck{TaskID: push.TaskID, Accept: true}))
writeEnv(t, conn, proto.NewEnvelope("id6", proto.TypeTaskResult, proto.TaskResult{
TaskID: push.TaskID, Status: "success",
PublishedURL: "https://example.com/v/1",
FinishedAt: time.Now().UnixMilli(),
}))
// 浏览器侧应收到三条 task.status(dispatched/running/success),第三条为 success
bEvents := []string{}
var lastPayload struct {
Status string `json:"status"`
}
for i := 0; i < 3; i++ {
e := readEnv(t, browserConn)
bEvents = append(bEvents, e.Type)
_ = json.Unmarshal(e.Payload, &lastPayload)
}
t.Logf("browser events: %v (last=%s)", bEvents, lastPayload.Status)
if bEvents[0] != "task.status" || bEvents[1] != "task.status" || bEvents[2] != "task.status" {
t.Fatalf("browser should receive 3 task.status events, got %v", bEvents)
}
if lastPayload.Status != "success" {
t.Fatalf("last event should be success, got %s", lastPayload.Status)
}
// 落库校验(轮询至 success)
deadline := time.Now().Add(3 * time.Second)
for {
code, env = doReq(t, "GET", fmt.Sprintf("/api/v1/tasks/%d", taskRow.ID), access, nil)
if code != 200 {
t.Fatalf("task get failed: %d", code)
}
_ = json.Unmarshal(env.Data, &taskRow)
if taskRow.Status == "success" || time.Now().After(deadline) {
break
}
time.Sleep(20 * time.Millisecond)
}
if taskRow.Status != "success" {
t.Fatalf("expect success, got %s", taskRow.Status)
}
if !strings.Contains(taskRow.PublishedURLs, "https://example.com/v/1") {
t.Fatalf("publishedUrls not recorded: %s", taskRow.PublishedURLs)
}
}
+18
View File
@@ -0,0 +1,18 @@
package auth
import (
"time"
"github.com/golang-jwt/jwt/v5"
)
func jwtRegisteredExpired() jwt.RegisteredClaims {
return jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(-time.Hour)),
Issuer: "everypublish",
}
}
func signExpired(secret string, claims AccessClaims) (string, error) {
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
}
+66
View File
@@ -0,0 +1,66 @@
package auth
import (
"strings"
"testing"
)
func TestHashVerifyPassword(t *testing.T) {
hash, err := HashPassword("everypublish-2026")
if err != nil {
t.Fatalf("hash: %v", err)
}
if !strings.HasPrefix(hash, "$argon2id$") {
t.Fatalf("bad hash format: %s", hash)
}
ok, err := VerifyPassword("everypublish-2026", hash)
if err != nil || !ok {
t.Fatalf("verify should pass: ok=%v err=%v", ok, err)
}
ok, err = VerifyPassword("wrong-pass", hash)
if err != nil || ok {
t.Fatalf("verify should fail: ok=%v err=%v", ok, err)
}
}
func TestJWTAccessRoundTrip(t *testing.T) {
secret := "test-secret"
token, err := IssueAccess(secret, 42, 7, "user", "owner")
if err != nil {
t.Fatalf("issue: %v", err)
}
claims, err := ParseAccess(secret, token)
if err != nil {
t.Fatalf("parse: %v", err)
}
if claims.UID != 42 || claims.WorkspaceID != 7 || claims.MemberRole != "owner" {
t.Fatalf("claims mismatch: %+v", claims)
}
}
func TestJWTAccessExpired(t *testing.T) {
secret := "test-secret"
claims := AccessClaims{UID: 1, RegisteredClaims: jwtRegisteredExpired()}
token, err := signExpired(secret, claims)
if err != nil {
t.Fatalf("sign: %v", err)
}
if _, err = ParseAccess(secret, token); err == nil {
t.Fatal("expired token should fail")
}
}
func TestJWTRefreshRoundTrip(t *testing.T) {
secret := "test-secret"
token, err := IssueRefresh(secret, 9, "jti-1")
if err != nil {
t.Fatalf("issue: %v", err)
}
claims, err := ParseRefresh(secret, token)
if err != nil {
t.Fatalf("parse: %v", err)
}
if claims.UID != 9 || claims.JTI != "jti-1" {
t.Fatalf("claims mismatch: %+v", claims)
}
}
+135
View File
@@ -0,0 +1,135 @@
package auth
import (
"errors"
"time"
"github.com/golang-jwt/jwt/v5"
)
// 过期策略:Access 15 分钟 / Refresh 7 天(旋转)/ Invite 24 小时
const (
AccessTTL = 15 * time.Minute
RefreshTTL = 7 * 24 * time.Hour
InviteTTL = 24 * time.Hour
)
// AccessClaims 访问令牌声明
type AccessClaims struct {
UID uint64 `json:"uid"`
WorkspaceID uint64 `json:"wsid"`
Role string `json:"role"`
MemberRole string `json:"mrole"`
jwt.RegisteredClaims
}
// RefreshClaims 刷新令牌声明
type RefreshClaims struct {
UID uint64 `json:"uid"`
JTI string `json:"jti"`
jwt.RegisteredClaims
}
// InviteClaims 邀请令牌声明
type InviteClaims struct {
WorkspaceID uint64 `json:"wsid"`
Email string `json:"email"`
Role string `json:"role"`
jwt.RegisteredClaims
}
// IssueAccess 签发访问令牌
func IssueAccess(secret string, uid, wsid uint64, role, memberRole string) (string, error) {
now := time.Now()
claims := AccessClaims{
UID: uid, WorkspaceID: wsid, Role: role, MemberRole: memberRole,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(AccessTTL)),
Issuer: "everypublish",
},
}
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
}
// IssueRefresh 签发刷新令牌
func IssueRefresh(secret string, uid uint64, jti string) (string, error) {
now := time.Now()
claims := RefreshClaims{
UID: uid, JTI: jti,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(RefreshTTL)),
Issuer: "everypublish",
},
}
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
}
// IssueInvite 签发邀请令牌
func IssueInvite(secret string, wsid uint64, email, role string) (string, error) {
now := time.Now()
claims := InviteClaims{
WorkspaceID: wsid, Email: email, Role: role,
RegisteredClaims: jwt.RegisteredClaims{
IssuedAt: jwt.NewNumericDate(now),
ExpiresAt: jwt.NewNumericDate(now.Add(InviteTTL)),
Issuer: "everypublish",
},
}
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(secret))
}
// ParseAccess 解析访问令牌
func ParseAccess(secret, token string) (*AccessClaims, error) {
claims := &AccessClaims{}
parsed, err := jwt.ParseWithClaims(token, claims, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, errors.New("unexpected signing method")
}
return []byte(secret), nil
})
if err != nil {
return nil, err
}
if !parsed.Valid {
return nil, errors.New("invalid token")
}
return claims, nil
}
// ParseRefresh 解析刷新令牌
func ParseRefresh(secret, token string) (*RefreshClaims, error) {
claims := &RefreshClaims{}
parsed, err := jwt.ParseWithClaims(token, claims, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, errors.New("unexpected signing method")
}
return []byte(secret), nil
})
if err != nil {
return nil, err
}
if !parsed.Valid {
return nil, errors.New("invalid token")
}
return claims, nil
}
// ParseInvite 解析邀请令牌
func ParseInvite(secret, token string) (*InviteClaims, error) {
claims := &InviteClaims{}
parsed, err := jwt.ParseWithClaims(token, claims, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, errors.New("unexpected signing method")
}
return []byte(secret), nil
})
if err != nil {
return nil, err
}
if !parsed.Valid {
return nil, errors.New("invalid token")
}
return claims, nil
}
+56
View File
@@ -0,0 +1,56 @@
package auth
import (
"crypto/rand"
"crypto/subtle"
"encoding/base64"
"fmt"
"strings"
"golang.org/x/crypto/argon2"
)
// Argon2id 参数(OWASP 推荐基线)
const (
argonTime = 1
argonMemory = 64 * 1024
argonThreads = 4
argonKeyLen = 32
saltLen = 16
)
// HashPassword 生成 $argon2id$v=19$m=65536,t=1,p=4$salt$hash
func HashPassword(password string) (string, error) {
salt := make([]byte, saltLen)
if _, err := rand.Read(salt); err != nil {
return "", err
}
hash := argon2.IDKey([]byte(password), salt, argonTime, argonMemory, argonThreads, argonKeyLen)
b64Salt := base64.RawStdEncoding.EncodeToString(salt)
b64Hash := base64.RawStdEncoding.EncodeToString(hash)
return fmt.Sprintf("$argon2id$v=19$m=%d,t=%d,p=%d$%s$%s", argonMemory, argonTime, argonThreads, b64Salt, b64Hash), nil
}
// VerifyPassword 校验密码(常数时间比较)
func VerifyPassword(password, encoded string) (bool, error) {
parts := strings.Split(encoded, "$")
if len(parts) != 6 {
return false, fmt.Errorf("invalid hash format")
}
var memory uint32
var timeCost uint32
var threads uint8
if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &timeCost, &threads); err != nil {
return false, err
}
salt, err := base64.RawStdEncoding.DecodeString(parts[4])
if err != nil {
return false, err
}
want, err := base64.RawStdEncoding.DecodeString(parts[5])
if err != nil {
return false, err
}
got := argon2.IDKey([]byte(password), salt, timeCost, memory, threads, uint32(len(want)))
return subtle.ConstantTimeCompare(got, want) == 1, nil
}
+31
View File
@@ -0,0 +1,31 @@
package cache
import (
"context"
"time"
"github.com/redis/go-redis/v9"
"everypublish/server/internal/config"
)
// Redis 封装(D4 起 asynq 队列同源共用)
type Redis struct {
client *redis.Client
}
// New 建立 Redis 连接
func New(cfg *config.Config) *Redis {
client := redis.NewClient(&redis.Options{Addr: cfg.RedisAddr, Password: cfg.RedisPass, DB: 0})
return &Redis{client: client}
}
// Ping 探测
func (r *Redis) Ping() error {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
return r.client.Ping(ctx).Err()
}
// Client 暴露底层客户端
func (r *Redis) Client() *redis.Client { return r.client }
+32
View File
@@ -0,0 +1,32 @@
package cache
import (
"context"
"fmt"
"time"
)
// SaveRefresh 登记刷新令牌(jti 白名单,用于轮换与登出吊销)
func (r *Redis) SaveRefresh(uid uint64, jti string, ttl time.Duration) error {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
key := fmt.Sprintf("refresh:%d:%s", uid, jti)
return r.client.Set(ctx, key, "1", ttl).Err()
}
// ExistsRefresh 校验刷新令牌是否有效
func (r *Redis) ExistsRefresh(uid uint64, jti string) (bool, error) {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
key := fmt.Sprintf("refresh:%d:%s", uid, jti)
n, err := r.client.Exists(ctx, key).Result()
return n == 1, err
}
// DeleteRefresh 吊销刷新令牌
func (r *Redis) DeleteRefresh(uid uint64, jti string) error {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
key := fmt.Sprintf("refresh:%d:%s", uid, jti)
return r.client.Del(ctx, key).Err()
}
+61
View File
@@ -0,0 +1,61 @@
package config
import (
"fmt"
"os"
)
// Config 服务器配置(环境变量优先,.env 兜底)
type Config struct {
ServerAddr string
MySQLDSN string
RedisAddr string
RedisPass string
JWTSecret string
BaseURL string
PublicBaseURL string // 浏览器侧访问地址(邀请链接等)
StorageDir string
StaticDir string // 前端生产构建目录(空=不托管静态)
MaxUploadBytes int64 // multipart 上传体积上限
}
// Load 从环境变量加载
func Load() *Config {
jwt := getenv("JWT_SECRET", "dev-secret-change-me")
if jwt == "dev-secret-change-me" || jwt == "please-change-me-in-production" || len(jwt) < 16 {
// 生产安全红线:默认/过短密钥可被伪造 JWT(含 admin 提权)。
// 保留 dev 默认值以便本地开发,但在任何环境都打印醒目告警。
println("!! [security] JWT_SECRET 使用默认值或长度不足(<16)。生产环境必须设置强随机 JWT_SECRET,否则任何人都能伪造令牌(admin 提权)。")
}
return &Config{
ServerAddr: getenv("SERVER_ADDR", ":8090"),
MySQLDSN: getenv("MYSQL_DSN", "root:everypublish@tcp(127.0.0.1:3306)/everypublish?charset=utf8mb4&parseTime=True&loc=Local"),
RedisAddr: getenv("REDIS_ADDR", "127.0.0.1:6379"),
RedisPass: getenv("REDIS_PASSWORD", ""),
JWTSecret: jwt,
BaseURL: getenv("BASE_URL", "http://127.0.0.1:8090"),
PublicBaseURL: getenv("PUBLIC_BASE_URL", getenv("BASE_URL", "http://127.0.0.1:8090")),
StorageDir: getenv("STORAGE_DIR", "./data/materials"),
StaticDir: getenv("STATIC_DIR", "../apps/web/dist"),
MaxUploadBytes: getenvInt64("MAX_UPLOAD_BYTES", 2<<30),
}
}
func getenvInt64(k string, def int64) int64 {
v := os.Getenv(k)
if v == "" {
return def
}
var n int64
if _, err := fmt.Sscanf(v, "%d", &n); err != nil || n <= 0 {
return def
}
return n
}
func getenv(k, def string) string {
if v := os.Getenv(k); v != "" {
return v
}
return def
}
+53
View File
@@ -0,0 +1,53 @@
package db
import (
"fmt"
"time"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
gormmysql "gorm.io/driver/mysql"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// Connect 建立 MySQL 连接并 Ping(重试 10 次,适配容器冷启动)
func Connect(cfg *config.Config) (*gorm.DB, error) {
db, err := gorm.Open(gormmysql.Open(cfg.MySQLDSN), &gorm.Config{Logger: logger.Default.LogMode(logger.Warn)})
if err != nil {
return nil, err
}
sqlDB, err := db.DB()
if err != nil {
return nil, err
}
sqlDB.SetMaxOpenConns(50)
sqlDB.SetMaxIdleConns(10)
sqlDB.SetConnMaxLifetime(time.Hour)
for i := 0; i < 10; i++ {
if err = sqlDB.Ping(); err == nil {
return db, nil
}
time.Sleep(time.Second)
}
return nil, fmt.Errorf("mysql ping failed: %w", err)
}
// AutoMigrate 开发期建表;生产走 golang-migrate(db/migrations)
func AutoMigrate(db *gorm.DB) error {
return db.AutoMigrate(
&models.User{},
&models.Workspace{},
&models.Member{},
&models.Account{},
&models.Material{},
&models.Task{},
&models.AgentDevice{},
&models.Challenge{},
&models.AuditLog{},
&models.Notification{},
&models.PairingCode{},
&models.TransferToken{},
)
}
+166
View File
@@ -0,0 +1,166 @@
package models
import "time"
// Base 通用字段
type Base struct {
ID uint64 `gorm:"primaryKey;autoIncrement" json:"id"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
}
// User 网站用户
type User struct {
Base
Email string `gorm:"uniqueIndex;size:191" json:"email"`
PasswordHash string `gorm:"size:255" json:"-"`
Nickname string `gorm:"size:64" json:"nickname"`
Role string `gorm:"size:16;default:user" json:"role"`
TOTPSecret string `gorm:"size:64" json:"-"`
TOTPEnabled bool `gorm:"default:false" json:"totpEnabled"`
Status string `gorm:"size:16;default:active" json:"status"`
}
// Workspace 工作空间(租户)
type Workspace struct {
Base
Name string `gorm:"size:128" json:"name"`
OwnerID uint64 `gorm:"index" json:"ownerId"`
Plan string `gorm:"size:32;default:free" json:"plan"`
Status string `gorm:"size:16;default:active" json:"status"`
}
// Member 工作空间成员(角色:owner/admin/operator/reviewer/viewer)
type Member struct {
Base
WorkspaceID uint64 `gorm:"uniqueIndex:uk_ws_user" json:"workspaceId"`
UserID uint64 `gorm:"uniqueIndex:uk_ws_user" json:"userId"`
Role string `gorm:"size:16;default:viewer" json:"role"`
}
// Account 平台账号(凭据仅在 Agent 本机保险库,服务器零凭据)
type Account struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Platform string `gorm:"size:32;index" json:"platform"`
AccountName string `gorm:"size:128" json:"accountName"`
Remark string `gorm:"size:255" json:"remark"`
AvatarURL string `gorm:"size:512" json:"avatarUrl"`
Status string `gorm:"size:16;default:unbound" json:"status"`
Health int `gorm:"default:100" json:"health"`
AgentDeviceID uint64 `gorm:"index" json:"agentDeviceId"`
IPProfile string `gorm:"size:191" json:"ipProfile"`
LastActiveAt *time.Time `json:"lastActiveAt"`
LastError string `gorm:"size:512" json:"lastError"`
}
// Material 素材
type Material struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Name string `gorm:"size:255" json:"name"`
Kind string `gorm:"size:16;default:video" json:"kind"`
Size int64 `json:"size"`
SHA256 string `gorm:"size:64;index" json:"sha256"`
Mime string `gorm:"size:128" json:"mime"`
StorageKey string `gorm:"size:512" json:"-"`
Group string `gorm:"size:64" json:"group"`
Tags string `gorm:"type:text" json:"tags"`
Status string `gorm:"size:16;default:ready" json:"status"`
CreatedBy uint64 `json:"createdBy"`
}
// Task 发布任务(状态机:draft→pending_review→approved→queued→dispatched→running→success/failed→retrying;可 cancelled/suspended/rejected)
type Task struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Title string `gorm:"size:255" json:"title"`
Content string `gorm:"type:text" json:"content"`
Tags string `gorm:"type:text" json:"tags"`
AccountIDs string `gorm:"type:text" json:"accountIds"`
MaterialIDs string `gorm:"type:text" json:"materialIds"`
ScheduleAt *time.Time `gorm:"index" json:"scheduleAt"`
Status string `gorm:"size:32;index;default:draft" json:"status"`
Priority int `gorm:"default:5" json:"priority"`
RetryCount int `gorm:"default:0" json:"retryCount"`
MaxRetry int `gorm:"default:3" json:"maxRetry"`
ErrorMessage string `gorm:"size:1024" json:"errorMessage"`
PublishedAt *time.Time `json:"publishedAt"`
PublishedURLs string `gorm:"type:text" json:"publishedUrls"`
Receipts string `gorm:"type:text" json:"receipts"`
CreatedBy uint64 `json:"createdBy"`
ReviewedBy uint64 `json:"reviewedBy"`
ReviewNote string `gorm:"size:512" json:"reviewNote"`
}
// AgentDevice 客户端设备(配对时注册 Ed25519 公钥)
type AgentDevice struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Name string `gorm:"size:128" json:"name"`
OS string `gorm:"size:32" json:"os"`
Version string `gorm:"size:32" json:"version"`
PublicKey string `gorm:"type:text" json:"-"`
Status string `gorm:"size:16;default:offline" json:"status"`
IP string `gorm:"size:64" json:"ip"`
Revoked bool `gorm:"default:false" json:"revoked"`
LastSeenAt *time.Time `json:"lastSeenAt"`
PairedAt time.Time `json:"pairedAt"`
}
// Challenge 二次验证挑战(qr/sms/confirm/captcha/pending;超时挂起,永不暴力破解)
type Challenge struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
AccountID uint64 `gorm:"index" json:"accountId"`
DeviceID uint64 `gorm:"index" json:"deviceId"`
Platform string `gorm:"size:32" json:"platform"`
Kind string `gorm:"size:16" json:"kind"`
Status string `gorm:"size:16;default:active" json:"status"`
QRToken string `gorm:"size:512" json:"qrToken"`
QRURL string `gorm:"size:1024" json:"qrUrl"`
Prompt string `gorm:"size:512" json:"prompt"`
Payload string `gorm:"type:text" json:"payload"`
ExpiresAt time.Time `gorm:"index" json:"expiresAt"`
SolvedAt *time.Time `json:"solvedAt"`
}
// AuditLog 审计日志
type AuditLog struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
UserID uint64 `gorm:"index" json:"userId"`
Action string `gorm:"size:64;index" json:"action"`
Resource string `gorm:"size:128" json:"resource"`
Detail string `gorm:"type:text" json:"detail"`
IP string `gorm:"size:64" json:"ip"`
}
// Notification 站内通知
type Notification struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
UserID uint64 `gorm:"index" json:"userId"`
Kind string `gorm:"size:32" json:"kind"`
Title string `gorm:"size:255" json:"title"`
Content string `gorm:"type:text" json:"content"`
Read bool `gorm:"default:false" json:"read"`
}
// PairingCode 配对码(一次性,5 分钟过期)
type PairingCode struct {
Base
WorkspaceID uint64 `gorm:"index" json:"workspaceId"`
Code string `gorm:"size:12;uniqueIndex" json:"code"`
ExpiresAt time.Time `json:"expiresAt"`
Used bool `gorm:"default:false" json:"used"`
DeviceID uint64 `json:"deviceId"`
}
// TransferToken 素材直链下载 token(短时签名)
type TransferToken struct {
Base
MaterialID uint64 `gorm:"index" json:"materialId"`
Token string `gorm:"size:191;uniqueIndex" json:"token"`
ExpiresAt time.Time `json:"expiresAt"`
}
+87
View File
@@ -0,0 +1,87 @@
package task
// 任务状态
const (
Draft = "draft"
PendingReview = "pending_review"
Rejected = "rejected"
Queued = "queued"
Dispatched = "dispatched"
Running = "running"
Success = "success"
Failed = "failed"
Suspended = "suspended"
Cancelled = "cancelled"
)
// 动作
type Action string
const (
ActionSubmit Action = "submit"
ActionApprove Action = "approve"
ActionReject Action = "reject"
ActionResubmit Action = "resubmit"
ActionDispatch Action = "dispatch"
ActionStart Action = "start"
ActionSuccess Action = "success"
ActionFail Action = "fail"
ActionRetry Action = "retry"
ActionSuspend Action = "suspend"
ActionResume Action = "resume"
ActionCancel Action = "cancel"
)
// transitions 状态转移表
var transitions = map[string]map[Action]string{
Draft: {
ActionSubmit: PendingReview,
ActionCancel: Cancelled,
},
PendingReview: {
ActionApprove: Queued,
ActionReject: Rejected,
ActionCancel: Cancelled,
},
Rejected: {
ActionResubmit: PendingReview,
ActionCancel: Cancelled,
},
Queued: {
ActionDispatch: Dispatched,
ActionCancel: Cancelled,
},
Dispatched: {
ActionStart: Running,
ActionFail: Failed,
ActionCancel: Cancelled,
},
Running: {
ActionSuccess: Success,
ActionFail: Failed,
ActionSuspend: Suspended,
ActionCancel: Cancelled,
},
Failed: {
ActionRetry: Queued,
ActionCancel: Cancelled,
},
Suspended: {
ActionResume: Queued,
ActionCancel: Cancelled,
},
Success: {},
Cancelled: {},
}
// Next 校验并返回下一状态;非法转移返回 ok=false
func Next(current string, a Action) (string, bool) {
next, ok := transitions[current][a]
return next, ok
}
// Can 当前状态是否接受该动作
func Can(current string, a Action) bool {
_, ok := transitions[current][a]
return ok
}
+82
View File
@@ -0,0 +1,82 @@
package task
import "testing"
func TestHappyPath(t *testing.T) {
path := []struct {
action Action
want string
}{
{ActionSubmit, PendingReview},
{ActionApprove, Queued},
{ActionDispatch, Dispatched},
{ActionStart, Running},
{ActionSuccess, Success},
}
cur := Draft
for _, step := range path {
next, ok := Next(cur, step.action)
if !ok || next != step.want {
t.Fatalf("step %s from %s: got %s ok=%v, want %s", step.action, cur, next, ok, step.want)
}
cur = next
}
}
func TestFailureRetry(t *testing.T) {
cur, ok := Next(Draft, ActionSubmit)
if !ok {
t.Fatal("submit failed")
}
cur, ok = Next(cur, ActionApprove)
if !ok {
t.Fatal("approve failed")
}
cur, ok = Next(cur, ActionDispatch)
if !ok {
t.Fatal("dispatch failed")
}
cur, ok = Next(cur, ActionStart)
if !ok {
t.Fatal("start failed")
}
cur, ok = Next(cur, ActionFail)
if !ok || cur != Failed {
t.Fatalf("fail: got %s ok=%v", cur, ok)
}
cur, ok = Next(cur, ActionRetry)
if !ok || cur != Queued {
t.Fatalf("retry: got %s ok=%v", cur, ok)
}
}
func TestSuspendResume(t *testing.T) {
cur, _ := Next(Running, ActionSuspend)
if cur != Suspended {
t.Fatalf("suspend: got %s", cur)
}
cur, ok := Next(cur, ActionResume)
if !ok || cur != Queued {
t.Fatalf("resume: got %s ok=%v", cur, ok)
}
}
func TestIllegalTransitions(t *testing.T) {
cases := []struct {
cur string
action Action
}{
{Draft, ActionApprove},
{PendingReview, ActionStart},
{Queued, ActionSuccess},
{Success, ActionRetry},
{Cancelled, ActionResubmit},
{Running, ActionApprove},
{Dispatched, ActionSuspend},
}
for _, c := range cases {
if _, ok := Next(c.cur, c.action); ok {
t.Fatalf("illegal transition should fail: %s -> %s", c.cur, c.action)
}
}
}
+250
View File
@@ -0,0 +1,250 @@
package ws
import (
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/hex"
"encoding/json"
"log"
"strconv"
"time"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
"everypublish/server/internal/task"
"everypublish/shared/proto"
)
// signPayload 签名原文:nonce + "|" + ts(两端一致)
func signPayload(nonce string, ts int64) []byte {
return []byte(nonce + "|" + strconv.FormatInt(ts, 10))
}
// HandleAgent 处理 Agent WSS 连接:hello 验签 → 注册 → 消息循环
func HandleAgent(cfg *config.Config, db *gorm.DB, hub *Hub, serverKey ed25519.PrivateKey) gin.HandlerFunc {
return func(c *gin.Context) {
conn, err := websocket.Accept(c.Writer, c.Request, &websocket.AcceptOptions{InsecureSkipVerify: true})
if err != nil {
return
}
// 首条消息必须是 hello(10s 超时)
ctx, cancel := context.WithTimeout(c.Request.Context(), 10*time.Second)
defer cancel()
typ, raw, err := conn.Read(ctx)
if err != nil {
_ = conn.Close(websocket.StatusPolicyViolation, "hello timeout")
return
}
if typ != websocket.MessageText {
_ = conn.Close(websocket.StatusPolicyViolation, "expect text")
return
}
var helloEnv proto.Envelope
if err = json.Unmarshal(raw, &helloEnv); err != nil || helloEnv.Type != proto.TypeHello {
_ = conn.Close(websocket.StatusPolicyViolation, "expect hello")
return
}
var hello proto.DeviceHello
if err = json.Unmarshal(helloEnv.Payload, &hello); err != nil {
_ = conn.Close(websocket.StatusPolicyViolation, "bad hello")
return
}
deviceID, _ := strconv.ParseUint(hello.DeviceID, 10, 64)
var device models.AgentDevice
if err = db.First(&device, deviceID).Error; err != nil {
_ = conn.Close(websocket.StatusPolicyViolation, "unknown device")
return
}
if device.Revoked {
_ = conn.Close(websocket.StatusPolicyViolation, "device revoked")
return
}
pubRaw, _ := base64.StdEncoding.DecodeString(device.PublicKey)
if len(pubRaw) != ed25519.PublicKeySize {
_ = conn.Close(websocket.StatusPolicyViolation, "bad key")
return
}
sig, _ := base64.StdEncoding.DecodeString(hello.Sig)
if !ed25519.Verify(ed25519.PublicKey(pubRaw), signPayload(hello.Nonce, hello.TS), sig) {
_ = conn.Close(websocket.StatusPolicyViolation, "bad signature")
return
}
// 注册 + 应答
serverNonce := randHex(16)
serverTS := time.Now().UnixMilli()
ackSig := ed25519.Sign(serverKey, signPayload(serverNonce, serverTS))
sess := &Conn{
ID: hello.DeviceID,
Kind: "agent",
WSID: device.WorkspaceID,
ws: conn,
send: make(chan []byte, 64),
}
hub.RegisterAgent(sess)
now := time.Now()
_ = db.Model(&device).Updates(map[string]interface{}{"status": "online", "last_seen_at": &now, "ip": c.ClientIP(), "version": hello.Version}).Error
hub.BroadcastWS(device.WorkspaceID, "agent.status", gin.H{"deviceId": device.ID, "online": true})
ack := proto.NewEnvelope(randHex(8), proto.TypeHelloAck, proto.HelloAck{
ServerNonce: serverNonce, SessionID: randHex(16), ServerTS: serverTS, Sig: base64.StdEncoding.EncodeToString(ackSig),
})
go sess.writeLoop()
sess.sendEnvelope(ack)
log.Printf("[ws] agent %d online (workspace %d)", device.ID, device.WorkspaceID)
// 消息循环
defer func() {
hub.UnregisterAgent(sess)
_ = db.Model(&device).Updates(map[string]interface{}{"status": "offline", "last_seen_at": time.Now()}).Error
hub.BroadcastWS(device.WorkspaceID, "agent.status", gin.H{"deviceId": device.ID, "online": false})
log.Printf("[ws] agent %d offline", device.ID)
}()
for {
// 90s 无消息视为死连接(客户端心跳 10-30s)
readCtx, readCancel := context.WithTimeout(c.Request.Context(), 90*time.Second)
typ, raw, err = conn.Read(readCtx)
readCancel()
if err != nil {
return
}
if typ != websocket.MessageText {
continue
}
var env proto.Envelope
if err = json.Unmarshal(raw, &env); err != nil {
continue
}
handleAgentMessage(db, hub, sess, device, env)
}
}
}
// handleAgentMessage 分发 Agent 消息(所有回写经 sess 统一通道,避免并发写)
func handleAgentMessage(db *gorm.DB, hub *Hub, sess *Conn, device models.AgentDevice, env proto.Envelope) {
var err error
switch env.Type {
case proto.TypeHeartbeat:
var hb proto.Heartbeat
if err = json.Unmarshal(env.Payload, &hb); err == nil {
now := time.Now()
_ = db.Model(&models.AgentDevice{}).Where("id = ?", device.ID).Update("last_seen_at", &now).Error
ack := proto.NewEnvelope(env.ID, proto.TypeHeartbeatAck, proto.HeartbeatAck{ServerTS: now.UnixMilli()})
sess.sendEnvelope(ack)
}
case proto.TypeTaskAck:
var ack proto.TaskAck
if err = json.Unmarshal(env.Payload, &ack); err == nil {
applyTaskAck(db, hub, device, ack)
}
case proto.TypeTaskResult:
var result proto.TaskResult
if err = json.Unmarshal(env.Payload, &result); err == nil {
log.Printf("[ws] task.result received: %+v", result)
applyTaskResult(db, hub, device, result)
} else {
log.Printf("[ws] task.result unmarshal error: %v payload=%s", err, string(env.Payload))
}
case proto.TypeChallengeAck:
var cack proto.ChallengeAck
if err = json.Unmarshal(env.Payload, &cack); err == nil {
log.Printf("[ws] challenge %s ack: %s", cack.ChallengeID, cack.Action)
}
case proto.TypeChallengeSolve:
var cs proto.ChallengeSolve
if err = json.Unmarshal(env.Payload, &cs); err == nil {
applyChallengeSolve(db, hub, device, cs)
}
default:
log.Printf("[ws] unknown message type: %s", env.Type)
}
}
// applyTaskAck 任务确认:dispatched → running
func applyTaskAck(db *gorm.DB, hub *Hub, device models.AgentDevice, ack proto.TaskAck) {
tid, _ := strconv.ParseUint(ack.TaskID, 10, 64)
var row models.Task
if err := db.Where("id = ? and workspace_id = ?", tid, device.WorkspaceID).First(&row).Error; err != nil {
return
}
if !ack.Accept {
_ = db.Model(&row).Updates(map[string]interface{}{"status": task.Failed, "error_message": ack.Reason}).Error
hub.BroadcastWS(device.WorkspaceID, "task.status", taskEvent(row))
return
}
if next, ok := task.Next(row.Status, task.ActionStart); ok {
_ = db.Model(&row).Update("status", next).Error
row.Status = next
hub.BroadcastWS(device.WorkspaceID, "task.status", taskEvent(row))
}
}
// applyTaskResult 结果回传:running → success/failed
func applyTaskResult(db *gorm.DB, hub *Hub, device models.AgentDevice, result proto.TaskResult) {
tid, _ := strconv.ParseUint(result.TaskID, 10, 64)
var row models.Task
if err := db.Where("id = ? and workspace_id = ?", tid, device.WorkspaceID).First(&row).Error; err != nil {
log.Printf("[ws] task.result: task %d not found: %v", tid, err)
return
}
log.Printf("[ws] task.result: task %d current status %s, want %s", tid, row.Status, result.Status)
if result.Status == "success" {
if next, ok := task.Next(row.Status, task.ActionSuccess); ok {
urls, _ := json.Marshal([]string{result.PublishedURL})
receipts, _ := json.Marshal(result.Receipts)
_ = db.Model(&row).Updates(map[string]interface{}{
"status": next, "published_urls": string(urls), "receipts": string(receipts),
"error_message": "", "published_at": time.Now(),
}).Error
row.Status = next
hub.BroadcastWS(device.WorkspaceID, "task.status", taskEvent(row))
log.Printf("[ws] task %d success in workspace %d", tid, device.WorkspaceID)
}
} else {
if next, ok := task.Next(row.Status, task.ActionFail); ok {
_ = db.Model(&row).Updates(map[string]interface{}{"status": next, "error_message": result.Error}).Error
row.Status = next
hub.BroadcastWS(device.WorkspaceID, "task.status", taskEvent(row))
}
}
}
// applyChallengeSolve 挑战完成:账号激活
func applyChallengeSolve(db *gorm.DB, hub *Hub, device models.AgentDevice, cs proto.ChallengeSolve) {
cid, _ := strconv.ParseUint(cs.ChallengeID, 10, 64)
var ch models.Challenge
if err := db.Where("id = ? and workspace_id = ?", cid, device.WorkspaceID).First(&ch).Error; err != nil {
return
}
now := time.Now()
_ = db.Model(&ch).Updates(map[string]interface{}{"status": "solved", "solved_at": &now, "qr_token": cs.Value}).Error
_ = db.Model(&models.Account{}).Where("id = ?", ch.AccountID).Updates(map[string]interface{}{"status": "active", "agent_device_id": device.ID, "last_active_at": &now}).Error
hub.BroadcastWS(device.WorkspaceID, "challenge.status", gin.H{"challengeId": ch.ID, "status": "solved", "accountId": ch.AccountID})
log.Printf("[ws] challenge %d solved, account %d active", ch.ID, ch.AccountID)
}
// taskEvent 任务状态事件
func taskEvent(row models.Task) gin.H {
return gin.H{"taskId": row.ID, "status": row.Status, "errorMessage": row.ErrorMessage, "publishedUrls": row.PublishedURLs}
}
// writeLoop 写协程(唯一写出口)
func (c *Conn) writeLoop() {
for raw := range c.send {
wctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
err := c.ws.Write(wctx, websocket.MessageText, raw)
cancel()
if err != nil {
return
}
}
}
func randHex(n int) string {
b := make([]byte, n)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
+45
View File
@@ -0,0 +1,45 @@
package ws
import (
"log"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"everypublish/server/internal/auth"
"everypublish/server/internal/config"
)
// HandleBrowser 浏览器实时通道:?token=JWT 鉴权,只读订阅工作空间事件
func HandleBrowser(cfg *config.Config, hub *Hub) gin.HandlerFunc {
return func(c *gin.Context) {
claims, err := auth.ParseAccess(cfg.JWTSecret, c.Query("token"))
if err != nil {
c.JSON(401, gin.H{"code": 1002, "message": "令牌无效"})
return
}
conn, err := websocket.Accept(c.Writer, c.Request, &websocket.AcceptOptions{InsecureSkipVerify: true})
if err != nil {
return
}
sess := &Conn{
ID: randHex(16),
Kind: "browser",
WSID: claims.WorkspaceID,
ws: conn,
send: make(chan []byte, 64),
}
hub.RegisterBrowser(sess)
go sess.writeLoop()
log.Printf("[ws] browser session %s online (workspace %d)", sess.ID, claims.WorkspaceID)
defer func() {
hub.UnregisterBrowser(sess)
log.Printf("[ws] browser session %s offline", sess.ID)
}()
for {
if _, _, err = conn.Read(c.Request.Context()); err != nil {
return
}
}
}
}
+161
View File
@@ -0,0 +1,161 @@
package ws
import (
"encoding/json"
"log"
"strconv"
"time"
"github.com/google/uuid"
"gorm.io/gorm"
"everypublish/server/internal/config"
"everypublish/server/internal/models"
"everypublish/server/internal/task"
"everypublish/shared/proto"
)
// Dispatcher 任务下发器(MVP:进程内轮询 + 审批/重试即时触发;接口预留二期换 asynq)
type Dispatcher struct {
DB *gorm.DB
Hub *Hub
Cfg *config.Config
Interval time.Duration
}
// NewDispatcher 构造下发器
func NewDispatcher(db *gorm.DB, hub *Hub, cfg *config.Config) *Dispatcher {
return &Dispatcher{DB: db, Hub: hub, Cfg: cfg, Interval: time.Second}
}
// Start 后台轮询:把到期且排队的任务下发给在线设备
func (d *Dispatcher) Start(stop <-chan struct{}) {
if d.Interval <= 0 {
d.Interval = time.Second
}
ticker := time.NewTicker(d.Interval)
defer ticker.Stop()
for {
select {
case <-stop:
return
case <-ticker.C:
d.poll()
}
}
}
func (d *Dispatcher) poll() {
var rows []models.Task
d.DB.Where("status = ?", task.Queued).
Where("(schedule_at is null or schedule_at <= ?)", time.Now()).
Order("priority desc, id asc").Limit(20).Find(&rows)
for i := range rows {
if d.TryDispatch(rows[i].ID) {
log.Printf("[dispatch] task %d pushed", rows[i].ID)
}
}
}
// TryDispatch 立即下发单个任务(审批/重试后调用);无在线设备返回 false
func (d *Dispatcher) TryDispatch(taskID uint64) bool {
var row models.Task
if err := d.DB.First(&row, taskID).Error; err != nil || row.Status != task.Queued {
return false
}
if row.ScheduleAt != nil && row.ScheduleAt.After(time.Now()) {
return false
}
devID := d.pickDevice(&row)
if devID == 0 {
return false
}
push := d.buildPush(&row, devID)
if push == nil {
return false
}
// 先落库 dispatched,再推送,避免 ack 早于状态落库的竞态
if next, ok := task.Next(row.Status, task.ActionDispatch); ok {
_ = d.DB.Model(&row).Update("status", next).Error
row.Status = next
}
env := proto.NewEnvelope(uuid.NewString(), proto.TypeTaskPush, push)
if err := d.Hub.SendToAgent(devID, env); err != nil {
// 下发失败(设备离线/队列满)回退 queued,避免卡在 dispatched
_ = d.DB.Model(&row).Update("status", task.Queued).Error
return false
}
d.Hub.BroadcastWS(row.WorkspaceID, "task.status", taskEvent(row))
return true
}
// pickDevice 选设备:账号已绑定设备则必须在线上才下发(保持一账号一IP),
// 未绑定设备的账号才兜底任意在线设备。
func (d *Dispatcher) pickDevice(row *models.Task) uint64 {
var accountIDs []uint64
_ = json.Unmarshal([]byte(row.AccountIDs), &accountIDs)
if len(accountIDs) > 0 {
var accounts []models.Account
d.DB.Where("id in ? and workspace_id = ?", accountIDs, row.WorkspaceID).Find(&accounts)
hasAssigned := false
for _, a := range accounts {
if a.AgentDeviceID == 0 {
continue
}
hasAssigned = true
if d.Hub.AgentOnline(a.AgentDeviceID) {
return a.AgentDeviceID
}
}
if hasAssigned {
return 0 // 已绑定设备但离线:不下发,保持 IP 隔离
}
}
online := d.Hub.OnlineDevices(row.WorkspaceID)
if len(online) > 0 {
return online[0]
}
return 0
}
// buildPush 组装 TaskPush(素材生成一次性签名直链)
func (d *Dispatcher) buildPush(row *models.Task, devID uint64) *proto.TaskPush {
var materialIDs []uint64
_ = json.Unmarshal([]byte(row.MaterialIDs), &materialIDs)
urls := make([]string, 0, len(materialIDs))
for _, mid := range materialIDs {
var mat models.Material
if err := d.DB.First(&mat, mid).Error; err != nil {
continue
}
token := models.TransferToken{MaterialID: mat.ID, Token: uuid.NewString(), ExpiresAt: time.Now().Add(10 * time.Minute)}
if err := d.DB.Create(&token).Error; err != nil {
continue
}
urls = append(urls, d.Cfg.BaseURL+"/api/v1/files/"+token.Token)
}
var accountIDs []uint64
_ = json.Unmarshal([]byte(row.AccountIDs), &accountIDs)
var tags []string
_ = json.Unmarshal([]byte(row.Tags), &tags)
push := &proto.TaskPush{
TaskID: strconv.FormatUint(row.ID, 10),
Title: row.Title,
Content: row.Content,
Tags: tags,
MaterialURLs: urls,
Priority: row.Priority,
}
if len(accountIDs) > 0 {
var acc models.Account
if err := d.DB.First(&acc, accountIDs[0]).Error; err == nil {
push.Platform = acc.Platform
push.AccountID = strconv.FormatUint(acc.ID, 10)
push.AccountName = acc.AccountName
}
}
if row.ScheduleAt != nil {
push.ScheduleAt = row.ScheduleAt.UnixMilli()
}
return push
}
+229
View File
@@ -0,0 +1,229 @@
package ws
import (
"encoding/json"
"errors"
"strconv"
"sync"
"time"
"github.com/coder/websocket"
"github.com/gin-gonic/gin"
"everypublish/shared/proto"
)
// parseID 解析设备ID字符串(失败返回 0)
func parseID(s string) uint64 {
id, _ := strconv.ParseUint(s, 10, 64)
return id
}
// Conn 一条已认证连接(Agent 或浏览器)
type Conn struct {
ID string
Kind string
WSID uint64
ws *websocket.Conn
send chan []byte
closeOnce sync.Once
mu sync.Mutex
closed bool
}
// close 安全关闭发送通道(仅一次)。
// 关闭在 mu 保护下进行;所有发送侧经 sendRaw 先检查 closed,
// 从根本上避免「send on closed channel」panic(关闭与发送的竞态)。
func (c *Conn) close() {
c.closeOnce.Do(func() {
c.mu.Lock()
c.closed = true
close(c.send)
c.mu.Unlock()
})
}
// sendRaw 向发送通道投递:closed 直接返回 false 不恐慌;
// block=true 时队列满将阻塞最多 1s(可靠消息用),否则满即丢弃(广播用)。
func (c *Conn) sendRaw(raw []byte, block bool) bool {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
return false
}
if block {
select {
case c.send <- raw:
return true
case <-time.After(time.Second):
return false
}
}
select {
case c.send <- raw:
return true
default:
return false
}
}
// sendEnvelope 统一写出口:所有消息经 send 通道由 writeLoop 单协程写入(避免并发写)
func (c *Conn) sendEnvelope(msg *proto.Envelope) {
raw, err := json.Marshal(msg)
if err != nil {
return
}
c.sendRaw(raw, true)
}
// Hub 连接中枢(goroutine-safe)
type Hub struct {
mu sync.RWMutex
agents map[uint64]*Conn
browser map[string]*Conn
}
// NewHub 创建中枢
func NewHub() *Hub {
return &Hub{
agents: make(map[uint64]*Conn),
browser: make(map[string]*Conn),
}
}
// RegisterAgent 上线登记(同设备重连替换旧连接)
// 网络关闭放到锁外执行,避免阻塞中枢其它操作。
func (h *Hub) RegisterAgent(c *Conn) {
var old *Conn
h.mu.Lock()
if o, ok := h.agents[parseID(c.ID)]; ok {
old = o
}
h.agents[parseID(c.ID)] = c
h.mu.Unlock()
if old != nil {
old.close()
if old.ws != nil {
_ = old.ws.Close(websocket.StatusGoingAway, "replaced by new connection")
}
}
}
// UnregisterAgent 下线(仅当仍是当前连接时才移除,避免旧连接误删新连接)
func (h *Hub) UnregisterAgent(c *Conn) {
var toClose *Conn
h.mu.Lock()
if cur, ok := h.agents[parseID(c.ID)]; ok && cur == c {
delete(h.agents, parseID(c.ID))
toClose = cur
}
h.mu.Unlock()
if toClose != nil {
toClose.close()
}
}
// RegisterBrowser 浏览器连接登记
func (h *Hub) RegisterBrowser(c *Conn) {
h.mu.Lock()
h.browser[c.ID] = c
h.mu.Unlock()
}
// UnregisterBrowser 浏览器断开(指针比较,防串号)
func (h *Hub) UnregisterBrowser(c *Conn) {
var toClose *Conn
h.mu.Lock()
if cur, ok := h.browser[c.ID]; ok && cur == c {
delete(h.browser, c.ID)
toClose = cur
}
h.mu.Unlock()
if toClose != nil {
toClose.close()
}
}
// Clear 清空所有连接(测试/停机用)
func (h *Hub) Clear() {
h.mu.Lock()
for id, c := range h.agents {
c.close()
_ = c.ws.Close(websocket.StatusGoingAway, "hub clear")
delete(h.agents, id)
}
for id, c := range h.browser {
c.close()
_ = c.ws.Close(websocket.StatusGoingAway, "hub clear")
delete(h.browser, id)
}
h.mu.Unlock()
}
// SendToAgent 给设备发消息(1s 超时;队列满丢弃)
func (h *Hub) SendToAgent(deviceID uint64, msg *proto.Envelope) error {
h.mu.RLock()
c, ok := h.agents[deviceID]
h.mu.RUnlock()
if !ok {
return errOffline
}
raw, err := json.Marshal(msg)
if err != nil {
return err
}
if !c.sendRaw(raw, true) {
return errBusy
}
return nil
}
// KickAgent 强制断开设备连接(吊销/远程指令用)
func (h *Hub) KickAgent(deviceID uint64) {
h.mu.RLock()
c, ok := h.agents[deviceID]
h.mu.RUnlock()
if ok {
_ = c.ws.Close(websocket.StatusPolicyViolation, "device revoked")
}
}
// AgentOnline 设备是否在线
func (h *Hub) AgentOnline(deviceID uint64) bool {
h.mu.RLock()
defer h.mu.RUnlock()
_, ok := h.agents[deviceID]
return ok
}
// OnlineDevices 工作空间在线设备列表
func (h *Hub) OnlineDevices(wsid uint64) []uint64 {
h.mu.RLock()
defer h.mu.RUnlock()
ids := make([]uint64, 0)
for id, c := range h.agents {
if c.WSID == wsid {
ids = append(ids, id)
}
}
return ids
}
// BroadcastWS 向工作空间所有浏览器连接广播事件
func (h *Hub) BroadcastWS(wsid uint64, typ string, payload interface{}) {
raw, err := json.Marshal(gin.H{"type": typ, "payload": payload})
if err != nil {
return
}
h.mu.RLock()
defer h.mu.RUnlock()
for _, c := range h.browser {
if c.WSID != wsid {
continue
}
c.sendRaw(raw, false)
}
}
var errOffline = errors.New("device offline")
var errBusy = errors.New("device send queue full")
+44
View File
@@ -0,0 +1,44 @@
package ws
import (
"testing"
)
// TestConnSendAfterClose 验证:连接关闭后 sendRaw/sendEnvelope 不 panic(回归 send-on-closed-channel)。
func TestConnSendAfterClose(t *testing.T) {
c := &Conn{ID: "1", send: make(chan []byte, 1), Kind: "agent"}
c.close()
// 已关闭,发送应返回 false 而非 panic
if c.sendRaw([]byte("x"), true) {
t.Fatal("send to closed conn should return false")
}
if c.sendRaw([]byte("x"), false) {
t.Fatal("send to closed conn (drop) should return false")
}
}
// TestConnDoubleClose 双重 close 不应 panic(closeOnce 保护)
func TestConnDoubleClose(t *testing.T) {
c := &Conn{ID: "1", send: make(chan []byte, 1)}
c.close()
c.close() // 第二次应无效果
}
// TestRegisterReplaceOld 同设备重连替换旧连接:旧连接被 close,映射指向新连接,且不 panic
func TestRegisterReplace(t *testing.T) {
h := NewHub()
old := &Conn{ID: "7", send: make(chan []byte, 64)}
newc := &Conn{ID: "7", send: make(chan []byte, 64)}
h.RegisterAgent(old)
h.RegisterAgent(newc)
if !h.AgentOnline(7) {
t.Fatal("device 7 should be online via new conn")
}
h.UnregisterAgent(old) // 注销旧连接不应影响新连接
if !h.AgentOnline(7) {
t.Fatal("device 7 should still be online after old conn unregister")
}
if old.sendRaw([]byte("x"), false) {
t.Fatal("old conn should be closed")
}
}
+30
View File
@@ -0,0 +1,30 @@
package ws
import (
"crypto/ed25519"
"crypto/rand"
"encoding/hex"
"os"
"path/filepath"
)
// LoadOrCreateServerKey 加载/生成服务器 Ed25519 密钥(hex 文件)
func LoadOrCreateServerKey(path string) (ed25519.PrivateKey, error) {
if raw, err := os.ReadFile(path); err == nil {
seed, err := hex.DecodeString(string(raw))
if err == nil && len(seed) == ed25519.SeedSize {
return ed25519.NewKeyFromSeed(seed), nil
}
}
_, priv, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
return nil, err
}
if err = os.MkdirAll(filepath.Dir(path), 0o755); err != nil {
return nil, err
}
if err = os.WriteFile(path, []byte(hex.EncodeToString(priv.Seed())), 0o600); err != nil {
return nil, err
}
return priv, nil
}
+77
View File
@@ -0,0 +1,77 @@
package main
import (
"context"
"log"
"os"
"os/signal"
"path/filepath"
"strings"
"syscall"
"github.com/gin-gonic/gin"
"github.com/joho/godotenv"
"everypublish/server/internal/api"
"everypublish/server/internal/cache"
"everypublish/server/internal/config"
"everypublish/server/internal/db"
"everypublish/server/internal/models"
"everypublish/server/internal/ws"
)
func main() {
_ = godotenv.Load() // 本地开发读 .env,生产用环境变量
cfg := config.Load()
gdb, err := db.Connect(cfg)
if err != nil {
log.Fatalf("db connect failed: %v", err)
}
if err = db.AutoMigrate(gdb); err != nil {
log.Fatalf("auto migrate failed: %v", err)
}
// 启动自愈:上次异常退出残留的 online 状态重置为 offline(真实状态以连接为准)
_ = gdb.Model(&models.AgentDevice{}).Where("status = ?", "online").Update("status", "offline").Error
log.Println("mysql connected & migrated")
rds := cache.New(cfg)
if err = rds.Ping(); err != nil {
log.Fatalf("redis ping failed: %v", err)
}
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()
dsp := ws.NewDispatcher(gdb, hub, cfg)
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)
// 托管前端生产构建(STATIC_DIR 存在 index.html 时启用),SPA 回退
if cfg.StaticDir != "" {
if _, err := os.Stat(filepath.Join(cfg.StaticDir, "index.html")); err == nil {
router.Static("/assets", filepath.Join(cfg.StaticDir, "assets"))
router.StaticFile("/favicon.ico", filepath.Join(cfg.StaticDir, "favicon.ico"))
router.NoRoute(func(c *gin.Context) {
if strings.HasPrefix(c.Request.URL.Path, "/api/") {
c.JSON(404, gin.H{"code": 404, "message": "接口不存在"})
return
}
c.File(filepath.Join(cfg.StaticDir, "index.html"))
})
log.Printf("serving frontend from %s", cfg.StaticDir)
}
}
log.Printf("everypublish-server listening on %s", cfg.ServerAddr)
if err = router.Run(cfg.ServerAddr); err != nil {
log.Fatalf("server exited: %v", err)
}
}