package main import ( "bytes" "encoding/json" "flag" "fmt" "io" "math/rand" "net/http" "os" "sort" "sync" "sync/atomic" "time" ) // 改进的压力测试工具 - 支持真实认证、梯度压测、详细统计 type TestConfig struct { BaseURL string Scenario string Duration time.Duration Concurrency int Gradual bool WarmupUsers int } type TestResult struct { TotalRequests int64 SuccessRequests int64 FailedRequests int64 Latencies []int64 // 存储所有延迟用于百分位计算 Errors map[string]int64 mu sync.Mutex } type AuthPool struct { tokens []string mu sync.RWMutex } type AdminToken struct { token string expiresAt time.Time mu sync.RWMutex } // 可用商品ID池 type ListingIDPool struct { ids []int mu sync.RWMutex } var globalListingIDs *ListingIDPool var ( baseURL = flag.String("url", "http://localhost:8080", "API 基础地址") concurrency = flag.Int("c", 50, "并发数") duration = flag.Int("d", 60, "测试时长(秒)") scenario = flag.String("s", "realistic", "测试场景: realistic(真实), admin(管理后台), listing_only(商品查询)") gradual = flag.Bool("gradual", false, "启用梯度压测") warmupUsers = flag.Int("warmup", 100, "预热用户数(生成token)") ) func main() { flag.Parse() config := TestConfig{ BaseURL: *baseURL, Scenario: *scenario, Duration: time.Duration(*duration) * time.Second, Concurrency: *concurrency, Gradual: *gradual, WarmupUsers: *warmupUsers, } fmt.Printf("=== 压力测试配置 ===\n") fmt.Printf("目标地址: %s\n", config.BaseURL) fmt.Printf("测试场景: %s\n", config.Scenario) fmt.Printf("测试时长: %d 秒\n", *duration) fmt.Printf("并发数: %d\n", config.Concurrency) fmt.Printf("梯度压测: %v\n", config.Gradual) fmt.Printf("==================\n\n") // 健康检查 if !healthCheck(config.BaseURL) { fmt.Println("❌ 后端服务未响应,退出测试") return } // 预加载可用商品ID fmt.Println("⏳ 预加载可用商品ID...") globalListingIDs = loadAvailableListings(config.BaseURL) if globalListingIDs == nil || len(globalListingIDs.ids) == 0 { fmt.Println("⚠️ 无法加载商品ID,将使用随机ID(可能导致404)") } else { fmt.Printf("✅ 成功加载 %d 个可用商品ID\n\n", len(globalListingIDs.ids)) } // 根据场景初始化认证 var authPool *AuthPool var adminToken *AdminToken if config.Scenario == "realistic" || config.Scenario == "listing_only" { fmt.Printf("⏳ 预热:生成 %d 个测试用户 token...\n", config.WarmupUsers) authPool = initAuthPool(config.BaseURL, config.WarmupUsers) if authPool == nil || len(authPool.tokens) == 0 { fmt.Println("⚠️ 无法生成用户token,将使用匿名访问") } else { fmt.Printf("✅ 成功生成 %d 个用户 token\n\n", len(authPool.tokens)) } } if config.Scenario == "admin" { fmt.Println("⏳ 获取管理员 token...") adminToken = initAdminToken(config.BaseURL) if adminToken == nil || adminToken.token == "" { fmt.Println("❌ 无法获取管理员token,退出测试") return } fmt.Println("✅ 成功获取管理员 token\n") } // 执行压测 var result *TestResult if config.Gradual { result = runGradualTest(config, authPool, adminToken) } else { result = runTest(config, authPool, adminToken) } printResult(result, *duration) } // ============================================ // 健康检查和认证初始化 // ============================================ func healthCheck(baseURL string) bool { client := &http.Client{Timeout: 5 * time.Second} resp, err := client.Get(baseURL + "/health") if err != nil { return false } defer resp.Body.Close() return resp.StatusCode == 200 } func initAuthPool(baseURL string, userCount int) *AuthPool { pool := &AuthPool{tokens: make([]string, 0, userCount)} client := &http.Client{Timeout: 10 * time.Second} successCount := 0 for i := 1; i <= userCount; i++ { phone := fmt.Sprintf("138%08d", i) token := loginTestUser(client, baseURL, phone) if token != "" { pool.tokens = append(pool.tokens, token) successCount++ } // 每10个打印一次进度 if i%10 == 0 || i == userCount { fmt.Printf("\r 进度: %d/%d (%d 成功)", i, userCount, successCount) } } fmt.Println() return pool } func loginTestUser(client *http.Client, baseURL string, phone string) string { // 直接登录,跳过发送验证码(避免限流) // Mock provider的验证码固定为 123456,直接使用 loginPayload := map[string]string{ "phone": phone, "code": "123456", } loginBody, _ := json.Marshal(loginPayload) resp, err := client.Post(baseURL+"/api/auth/sms/login", "application/json", bytes.NewBuffer(loginBody)) if err != nil { return "" } defer resp.Body.Close() // 如果验证码无效,发送一次验证码后重试 if resp.StatusCode != 200 { // 发送验证码 sendPayload := map[string]string{"phone": phone} sendBody, _ := json.Marshal(sendPayload) sendResp, err := client.Post(baseURL+"/api/auth/sms/send", "application/json", bytes.NewBuffer(sendBody)) if err != nil { return "" } io.Copy(io.Discard, sendResp.Body) sendResp.Body.Close() // 等待验证码写入Redis time.Sleep(200 * time.Millisecond) // 重试登录 resp2, err := client.Post(baseURL+"/api/auth/sms/login", "application/json", bytes.NewBuffer(loginBody)) if err != nil { return "" } defer resp2.Body.Close() if resp2.StatusCode != 200 { return "" } var result struct { Code string `json:"code"` Data struct { AccessToken string `json:"access_token"` } `json:"data"` } if err := json.NewDecoder(resp2.Body).Decode(&result); err != nil { return "" } if result.Code == "ok" { return result.Data.AccessToken } return "" } var result struct { Code string `json:"code"` // 修复:API返回字符串"ok" Data struct { AccessToken string `json:"access_token"` } `json:"data"` } if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { return "" } if result.Code == "ok" { return result.Data.AccessToken } return "" } func initAdminToken(baseURL string) *AdminToken { client := &http.Client{Timeout: 10 * time.Second} // 获取验证码(为了获取captcha_id) resp, err := client.Get(baseURL + "/api/admin/auth/captcha") if err != nil { fmt.Printf(" 错误: %v\n", err) return nil } defer resp.Body.Close() var captchaResp struct { Data struct { CaptchaID string `json:"captcha_id"` } `json:"data"` } if err := json.NewDecoder(resp.Body).Decode(&captchaResp); err != nil { return nil } adminUsername := os.Getenv("LOAD_STRESS_ADMIN_USERNAME") adminPassword := os.Getenv("LOAD_STRESS_ADMIN_PASSWORD") if adminUsername == "" || adminPassword == "" { fmt.Println(" 跳过后台登录:请设置 LOAD_STRESS_ADMIN_USERNAME 和 LOAD_STRESS_ADMIN_PASSWORD") return nil } // 登录后台账号 loginPayload := map[string]string{ "username": adminUsername, "password": adminPassword, "captcha_id": captchaResp.Data.CaptchaID, "captcha_code": "1234", } loginBody, _ := json.Marshal(loginPayload) resp, err = client.Post(baseURL+"/api/admin/auth/login", "application/json", bytes.NewBuffer(loginBody)) if err != nil { fmt.Printf(" 错误: %v\n", err) return nil } defer resp.Body.Close() if resp.StatusCode != 200 { fmt.Printf(" 登录失败: HTTP %d\n", resp.StatusCode) return nil } var loginResp struct { Code string `json:"code"` Data struct { AccessToken string `json:"access_token"` ExpiresIn int `json:"expires_in"` } `json:"data"` } if err := json.NewDecoder(resp.Body).Decode(&loginResp); err != nil { return nil } if loginResp.Code != "ok" { fmt.Printf(" 登录失败: code=%s\n", loginResp.Code) return nil } return &AdminToken{ token: loginResp.Data.AccessToken, expiresAt: time.Now().Add(time.Duration(loginResp.Data.ExpiresIn) * time.Second), } } func (p *AuthPool) GetRandomToken() string { if p == nil || len(p.tokens) == 0 { return "" } p.mu.RLock() defer p.mu.RUnlock() return p.tokens[rand.Intn(len(p.tokens))] } func (a *AdminToken) Get() string { if a == nil { return "" } a.mu.RLock() defer a.mu.RUnlock() return a.token } func (p *ListingIDPool) GetRandomID() int { if p == nil || len(p.ids) == 0 { // 降级:返回随机ID return rand.Intn(5000) + 1 } p.mu.RLock() defer p.mu.RUnlock() return p.ids[rand.Intn(len(p.ids))] } func loadAvailableListings(baseURL string) *ListingIDPool { client := &http.Client{Timeout: 10 * time.Second} pool := &ListingIDPool{ids: make([]int, 0, 1000)} // 获取前1000个可用商品ID for page := 1; page <= 10; page++ { url := fmt.Sprintf("%s/api/listings?page=%d&page_size=100", baseURL, page) resp, err := client.Get(url) if err != nil { break } var result struct { Code string `json:"code"` // 修复:API返回字符串"ok"而不是数字0 Data struct { Items []struct { ID int `json:"id"` } `json:"items"` } `json:"data"` } if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { resp.Body.Close() break } resp.Body.Close() if result.Code != "ok" || len(result.Data.Items) == 0 { break } for _, item := range result.Data.Items { pool.ids = append(pool.ids, item.ID) } // 如果不满100个,说明已经到最后一页 if len(result.Data.Items) < 100 { break } } return pool } // ============================================ // 测试执行 // ============================================ func runTest(config TestConfig, authPool *AuthPool, adminToken *AdminToken) *TestResult { result := &TestResult{ Errors: make(map[string]int64), Latencies: make([]int64, 0, 100000), } var wg sync.WaitGroup stopChan := make(chan struct{}) fmt.Printf("🚀 开始压测 (%d 并发, %v)...\n\n", config.Concurrency, config.Duration) // 启动统计goroutine statsTicker := time.NewTicker(5 * time.Second) go func() { for { select { case <-statsTicker.C: printProgress(result) case <-stopChan: statsTicker.Stop() return } } }() // 启动并发workers for i := 0; i < config.Concurrency; i++ { wg.Add(1) go func(workerID int) { defer wg.Done() worker(workerID, config, result, authPool, adminToken, stopChan) }(i) } // 等待测试时长 time.Sleep(config.Duration) close(stopChan) wg.Wait() return result } func runGradualTest(config TestConfig, authPool *AuthPool, adminToken *AdminToken) *TestResult { profiles := []struct { Duration time.Duration Concurrency int }{ {30 * time.Second, 10}, // 预热 {60 * time.Second, config.Concurrency / 4}, // 25%负载 {60 * time.Second, config.Concurrency / 2}, // 50%负载 {60 * time.Second, config.Concurrency}, // 100%负载 {30 * time.Second, config.Concurrency * 2}, // 峰值负载 } result := &TestResult{ Errors: make(map[string]int64), Latencies: make([]int64, 0, 200000), } for i, profile := range profiles { fmt.Printf("📊 阶段 %d/%d: %d 并发, 持续 %v\n", i+1, len(profiles), profile.Concurrency, profile.Duration) phaseConfig := config phaseConfig.Duration = profile.Duration phaseConfig.Concurrency = profile.Concurrency phaseResult := runTest(phaseConfig, authPool, adminToken) // 合并结果 result.mu.Lock() result.TotalRequests += phaseResult.TotalRequests result.SuccessRequests += phaseResult.SuccessRequests result.FailedRequests += phaseResult.FailedRequests result.Latencies = append(result.Latencies, phaseResult.Latencies...) for k, v := range phaseResult.Errors { result.Errors[k] += v } result.mu.Unlock() if i < len(profiles)-1 { fmt.Println("⏸ 冷却 10 秒...") time.Sleep(10 * time.Second) } } return result } func worker(id int, config TestConfig, result *TestResult, authPool *AuthPool, adminToken *AdminToken, stopChan chan struct{}) { client := &http.Client{ Timeout: 10 * time.Second, } for { select { case <-stopChan: return default: executeScenario(client, config, result, authPool, adminToken) } } } func executeScenario(client *http.Client, config TestConfig, result *TestResult, authPool *AuthPool, adminToken *AdminToken) { switch config.Scenario { case "realistic": executeRealisticScenario(client, config.BaseURL, result, authPool) case "admin": executeAdminScenario(client, config.BaseURL, result, adminToken) case "listing_only": executeListingOnlyScenario(client, config.BaseURL, result, authPool) default: testHealthCheck(client, config.BaseURL, result) } } // ============================================ // 业务场景实现 // ============================================ func executeRealisticScenario(client *http.Client, baseURL string, result *TestResult, authPool *AuthPool) { r := rand.Intn(1000) switch { case r < 350: // 35% 查询商品列表 testListListings(client, baseURL, result) case r < 600: // 25% 查询商品详情 testGetListingDetail(client, baseURL, result) case r < 750: // 15% 查询我的订单 testListMyOrders(client, baseURL, result, authPool) case r < 850: // 10% 查询钱包余额 testWalletBalance(client, baseURL, result, authPool) case r < 920: // 7% 聊天列表 testChatList(client, baseURL, result, authPool) case r < 970: // 5% 查询订单详情 testOrderDetail(client, baseURL, result, authPool) case r < 990: // 2% 创建订单(需要认证) testCreateOrder(client, baseURL, result, authPool) default: // 1% 支付订单(需要认证) testPayOrder(client, baseURL, result, authPool) } } func executeAdminScenario(client *http.Client, baseURL string, result *TestResult, adminToken *AdminToken) { r := rand.Intn(100) switch { case r < 25: // 25% 用户管理列表 testAdminUsers(client, baseURL, result, adminToken) case r < 45: // 20% 订单管理列表 testAdminOrders(client, baseURL, result, adminToken) case r < 60: // 15% 商品审核列表 testAdminListings(client, baseURL, result, adminToken) case r < 75: // 15% 钱包流水 testAdminWalletLedger(client, baseURL, result, adminToken) case r < 85: // 10% 申诉管理 testAdminDisputes(client, baseURL, result, adminToken) case r < 95: // 10% 审计日志 testAdminAuditLogs(client, baseURL, result, adminToken) default: // 5% 仪表盘 testAdminDashboard(client, baseURL, result, adminToken) } } func executeListingOnlyScenario(client *http.Client, baseURL string, result *TestResult, authPool *AuthPool) { r := rand.Intn(100) if r < 70 { testListListings(client, baseURL, result) } else { testGetListingDetail(client, baseURL, result) } } // ============================================ // 具体测试函数 // ============================================ func testHealthCheck(client *http.Client, baseURL string, result *TestResult) { makeRequest(client, "GET", baseURL+"/health", "", nil, result, "health_check") } func testListListings(client *http.Client, baseURL string, result *TestResult) { page := rand.Intn(10) + 1 pageSize := []int{10, 20, 50}[rand.Intn(3)] url := fmt.Sprintf("%s/api/listings?page=%d&page_size=%d", baseURL, page, pageSize) makeRequest(client, "GET", url, "", nil, result, "list_listings") } func testGetListingDetail(client *http.Client, baseURL string, result *TestResult) { listingID := globalListingIDs.GetRandomID() url := fmt.Sprintf("%s/api/listings/%d", baseURL, listingID) makeRequest(client, "GET", url, "", nil, result, "get_listing_detail") } func testListMyOrders(client *http.Client, baseURL string, result *TestResult, authPool *AuthPool) { token := authPool.GetRandomToken() if token == "" { return } page := rand.Intn(5) + 1 url := fmt.Sprintf("%s/api/orders?page=%d&page_size=20", baseURL, page) makeAuthRequest(client, "GET", url, token, nil, result, "list_my_orders") } func testWalletBalance(client *http.Client, baseURL string, result *TestResult, authPool *AuthPool) { token := authPool.GetRandomToken() if token == "" { return } makeAuthRequest(client, "GET", baseURL+"/api/wallet/balance", token, nil, result, "wallet_balance") } func testChatList(client *http.Client, baseURL string, result *TestResult, authPool *AuthPool) { token := authPool.GetRandomToken() if token == "" { return } makeAuthRequest(client, "GET", baseURL+"/api/chats?page=1&page_size=20", token, nil, result, "chat_list") } func testOrderDetail(client *http.Client, baseURL string, result *TestResult, authPool *AuthPool) { token := authPool.GetRandomToken() if token == "" { return } orderID := rand.Intn(3000) + 1 url := fmt.Sprintf("%s/api/orders/%d", baseURL, orderID) makeAuthRequest(client, "GET", url, token, nil, result, "order_detail") } func testCreateOrder(client *http.Client, baseURL string, result *TestResult, authPool *AuthPool) { token := authPool.GetRandomToken() if token == "" { return } listingID := globalListingIDs.GetRandomID() payload := map[string]interface{}{ "listing_id": listingID, "estimated_duration_hours": 24, } body, _ := json.Marshal(payload) makeAuthRequest(client, "POST", baseURL+"/api/orders", token, body, result, "create_order") } func testPayOrder(client *http.Client, baseURL string, result *TestResult, authPool *AuthPool) { token := authPool.GetRandomToken() if token == "" { return } orderID := rand.Intn(3000) + 1 url := fmt.Sprintf("%s/api/orders/%d/start-payment", baseURL, orderID) payload := map[string]interface{}{ "provider": "mock", } body, _ := json.Marshal(payload) makeAuthRequest(client, "POST", url, token, body, result, "pay_order") } // 管理后台测试函数 func testAdminUsers(client *http.Client, baseURL string, result *TestResult, adminToken *AdminToken) { page := rand.Intn(10) + 1 url := fmt.Sprintf("%s/api/admin/users?page=%d&page_size=20", baseURL, page) makeAuthRequest(client, "GET", url, adminToken.Get(), nil, result, "admin_users") } func testAdminOrders(client *http.Client, baseURL string, result *TestResult, adminToken *AdminToken) { page := rand.Intn(10) + 1 url := fmt.Sprintf("%s/api/admin/orders?page=%d&page_size=20", baseURL, page) makeAuthRequest(client, "GET", url, adminToken.Get(), nil, result, "admin_orders") } func testAdminListings(client *http.Client, baseURL string, result *TestResult, adminToken *AdminToken) { page := rand.Intn(10) + 1 url := fmt.Sprintf("%s/api/admin/listings?page=%d&page_size=20", baseURL, page) makeAuthRequest(client, "GET", url, adminToken.Get(), nil, result, "admin_listings") } func testAdminWalletLedger(client *http.Client, baseURL string, result *TestResult, adminToken *AdminToken) { page := rand.Intn(20) + 1 url := fmt.Sprintf("%s/api/admin/wallet/ledger?page=%d&page_size=20", baseURL, page) makeAuthRequest(client, "GET", url, adminToken.Get(), nil, result, "admin_wallet_ledger") } func testAdminDisputes(client *http.Client, baseURL string, result *TestResult, adminToken *AdminToken) { page := rand.Intn(5) + 1 url := fmt.Sprintf("%s/api/admin/disputes?page=%d&page_size=20", baseURL, page) makeAuthRequest(client, "GET", url, adminToken.Get(), nil, result, "admin_disputes") } func testAdminAuditLogs(client *http.Client, baseURL string, result *TestResult, adminToken *AdminToken) { page := rand.Intn(20) + 1 url := fmt.Sprintf("%s/api/admin/audit-logs?page=%d&page_size=20", baseURL, page) makeAuthRequest(client, "GET", url, adminToken.Get(), nil, result, "admin_audit_logs") } func testAdminDashboard(client *http.Client, baseURL string, result *TestResult, adminToken *AdminToken) { makeAuthRequest(client, "GET", baseURL+"/api/admin/dashboard", adminToken.Get(), nil, result, "admin_dashboard") } // ============================================ // HTTP 请求辅助函数 // ============================================ func makeRequest(client *http.Client, method, url, token string, body []byte, result *TestResult, apiName string) { var req *http.Request var err error if body != nil { req, err = http.NewRequest(method, url, bytes.NewBuffer(body)) } else { req, err = http.NewRequest(method, url, nil) } if err != nil { recordError(result, apiName+"_req_error") atomic.AddInt64(&result.TotalRequests, 1) atomic.AddInt64(&result.FailedRequests, 1) return } if token != "" { req.Header.Set("Authorization", "Bearer "+token) } if body != nil { req.Header.Set("Content-Type", "application/json") } start := time.Now() resp, err := client.Do(req) latency := time.Since(start).Milliseconds() atomic.AddInt64(&result.TotalRequests, 1) recordLatency(result, latency) if err != nil { atomic.AddInt64(&result.FailedRequests, 1) recordError(result, apiName+"_error: "+err.Error()) return } defer resp.Body.Close() io.Copy(io.Discard, resp.Body) if resp.StatusCode >= 200 && resp.StatusCode < 300 { atomic.AddInt64(&result.SuccessRequests, 1) } else { atomic.AddInt64(&result.FailedRequests, 1) recordError(result, fmt.Sprintf("%s_status_%d", apiName, resp.StatusCode)) } } func makeAuthRequest(client *http.Client, method, url, token string, body []byte, result *TestResult, apiName string) { makeRequest(client, method, url, token, body, result, apiName) } func recordLatency(result *TestResult, latency int64) { result.mu.Lock() defer result.mu.Unlock() result.Latencies = append(result.Latencies, latency) } func recordError(result *TestResult, errMsg string) { result.mu.Lock() defer result.mu.Unlock() result.Errors[errMsg]++ } // ============================================ // 结果统计和输出 // ============================================ func printProgress(result *TestResult) { total := atomic.LoadInt64(&result.TotalRequests) success := atomic.LoadInt64(&result.SuccessRequests) failed := atomic.LoadInt64(&result.FailedRequests) if total > 0 { successRate := float64(success) / float64(total) * 100 fmt.Printf(" 进行中: %d 请求 | 成功率: %.2f%% | 失败: %d\n", total, successRate, failed) } } func printResult(result *TestResult, durationSec int) { fmt.Printf("\n\n=== 压力测试结果 ===\n") fmt.Printf("总请求数: %d\n", result.TotalRequests) fmt.Printf("成功请求: %d (%.2f%%)\n", result.SuccessRequests, float64(result.SuccessRequests)/float64(result.TotalRequests)*100) fmt.Printf("失败请求: %d (%.2f%%)\n", result.FailedRequests, float64(result.FailedRequests)/float64(result.TotalRequests)*100) qps := float64(result.TotalRequests) / float64(durationSec) fmt.Printf("\nQPS: %.2f\n", qps) if len(result.Latencies) > 0 { result.mu.Lock() latencies := make([]int64, len(result.Latencies)) copy(latencies, result.Latencies) result.mu.Unlock() sort.Slice(latencies, func(i, j int) bool { return latencies[i] < latencies[j] }) p50 := latencies[len(latencies)*50/100] p95 := latencies[len(latencies)*95/100] p99 := latencies[len(latencies)*99/100] min := latencies[0] max := latencies[len(latencies)-1] var sum int64 for _, l := range latencies { sum += l } avg := sum / int64(len(latencies)) fmt.Printf("\n延迟统计:\n") fmt.Printf(" 最小: %d ms\n", min) fmt.Printf(" P50: %d ms\n", p50) fmt.Printf(" 平均: %d ms\n", avg) fmt.Printf(" P95: %d ms\n", p95) fmt.Printf(" P99: %d ms\n", p99) fmt.Printf(" 最大: %d ms\n", max) } if len(result.Errors) > 0 { fmt.Printf("\n错误统计 (Top 10):\n") type errorPair struct { msg string count int64 } var errors []errorPair for msg, count := range result.Errors { errors = append(errors, errorPair{msg, count}) } sort.Slice(errors, func(i, j int) bool { return errors[i].count > errors[j].count }) for i, e := range errors { if i >= 10 { break } fmt.Printf(" %s: %d 次\n", e.msg, e.count) } } fmt.Printf("==================\n") }