251 lines
7.9 KiB
Go
251 lines
7.9 KiB
Go
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)
|
||
}
|
||
}
|