Files
EveryPublish/server/internal/api/router_test.go
T

295 lines
9.8 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package api_test
import (
"bytes"
"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(1)
}
if err = adm.Exec("CREATE DATABASE IF NOT EXISTS everypublish_test CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci").Error; err != nil {
fmt.Println("skip: create test db failed:", err)
os.Exit(1)
}
testDSN := "root:everypublish@tcp(127.0.0.1:3306)/everypublish_test?charset=utf8mb4&parseTime=True&loc=Local"
testDB, err = gorm.Open(gormmysql.Open(testDSN), &gorm.Config{})
if err != nil {
fmt.Println("skip: test db connect failed:", err)
os.Exit(1)
}
resetTestDB()
cfg := &config.Config{
ServerAddr: ":0",
MySQLDSN: testDSN,
RedisAddr: "127.0.0.1:6379",
JWTSecret: "test-secret",
BaseURL: "http://127.0.0.1:8090",
StorageDir: "/tmp/everypublish-test",
ExecutorMode: "mock",
}
rds := cache.New(cfg)
testHub = ws.NewHub()
dsp := ws.NewDispatcher(testDB, testHub, cfg)
testRouter = api.Router(cfg, testDB, rds, testHub, dsp)
code := m.Run()
os.Exit(code)
}
// resetTestDB 清空并重建表 + 清空连接中枢(保证 -count 多轮可重复)
func resetTestDB() {
if testHub != nil {
testHub.Clear()
}
for _, t := range []interface{}{
&models.TransferToken{}, &models.Notification{}, &models.AuditLog{},
&models.Challenge{}, &models.AgentDevice{}, &models.PairingCode{}, &models.TaskEvent{}, &models.Task{}, &models.Material{},
&models.Account{}, &models.Member{}, &models.Workspace{}, &models.User{},
} {
testDB.Migrator().DropTable(t)
}
_ = 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)
}
var secondWS models.Workspace
_ = json.Unmarshal(env.Data, &secondWS)
code, _ = doReq(t, "PUT", fmt.Sprintf("/api/v1/workspaces/%d", secondWS.ID), aAccess, gin.H{"name": "越权空间"})
if code != 404 {
t.Fatalf("cross-workspace update should be 404, got %d", code)
}
var ownWorkspaces struct {
List []models.Workspace `json:"list"`
}
code, env = doReq(t, "GET", "/api/v1/workspaces", aAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("list own workspaces failed: %d %s", code, env.Message)
}
_ = json.Unmarshal(env.Data, &ownWorkspaces)
if len(ownWorkspaces.List) == 0 {
t.Fatal("owner should have at least one workspace")
}
primaryWSID := ownWorkspaces.List[0].ID
// A 邀请 B(viewer)
code, env = doReq(t, "POST", "/api/v1/members/invite", aAccess, gin.H{"email": "worker@test.com", "role": "viewer"})
if code != 200 || env.Code != 0 {
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)
}
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/auth/switch-workspace/%d", primaryWSID), bAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("workspace switch failed: %d %s", code, env.Message)
}
var switched tokens
_ = json.Unmarshal(env.Data, &switched)
bAccess = switched.AccessToken
code, _ = doReq(t, "POST", "/api/v1/accounts", bAccess, gin.H{"platform": "douyin", "remark": "viewer-forbidden"})
if code != 403 {
t.Fatalf("viewer account write should be 403, got %d", code)
}
// A 查成员:2 人
code, env = doReq(t, "GET", "/api/v1/members", aAccess, nil)
if code != 200 || env.Code != 0 {
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)
}
code, env = doReq(t, "POST", fmt.Sprintf("/api/v1/auth/switch-workspace/%d", primaryWSID), bAccess, nil)
if code != 200 || env.Code != 0 {
t.Fatalf("workspace switch after role update failed: %d %s", code, env.Message)
}
_ = json.Unmarshal(env.Data, &switched)
var switchInfo struct {
MemberRole string `json:"memberRole"`
}
_ = json.Unmarshal(env.Data, &switchInfo)
if switchInfo.MemberRole != "operator" {
t.Fatalf("workspace switch should reflect updated role, got %q", switchInfo.MemberRole)
}
bAccess = switched.AccessToken
code, env = doReq(t, "POST", "/api/v1/accounts", bAccess, gin.H{"platform": "douyin", "remark": "operator-allowed"})
if code != 200 || env.Code != 0 {
t.Fatalf("operator account write should be allowed: %d %s", code, env.Message)
}
// owner 不可被移除
code, _ = doReq(t, "DELETE", fmt.Sprintf("/api/v1/members/%d", ownerMember), aAccess, nil)
if code != 403 {
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)
}
}