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) } }