package hotime import ( "bytes" "fmt" "io" "net/http" "net/http/httptest" "os" "os/signal" "sort" "strings" "sync" "sync/atomic" "syscall" "testing" "time" . "code.hoteas.com/golang/hotime/db" ) // TestApp 测试应用,封装 Application 提供测试能力 type TestApp struct { *Application configPath string listeners []func(*Context) bool workDir string projs TestProj flows FlowTest collector *TestCollector swaggerEnabled bool swaggerTitle string swaggerVersion string swaggerDir string stopNewTests atomic.Bool swaggerMu sync.Mutex inited bool initOnce sync.Once } // Test 创建测试应用(链式入口,延迟 Init 以便先 WorkDir) func Test(configPath string, listeners ...func(*Context) bool) *TestApp { return &TestApp{ configPath: configPath, listeners: listeners, collector: &TestCollector{}, } } // WorkDir 设置工作目录(在 Init 前 chdir) func (that *TestApp) WorkDir(dir string) *TestApp { that.workDir = dir return that } // Swagger 启用接口级增量 Swagger 写入。 // 参数可选:Swagger() / Swagger(title) / Swagger(title, version) / Swagger(title, version, outputDir) func (that *TestApp) Swagger(args ...string) *TestApp { that.swaggerEnabled = true that.swaggerTitle = "API" that.swaggerVersion = "1.0.0" if len(args) > 0 && args[0] != "" { that.swaggerTitle = args[0] } if len(args) > 1 && args[1] != "" { that.swaggerVersion = args[1] } if len(args) > 2 && args[2] != "" { that.swaggerDir = args[2] } return that } // Proj 注册路由与单接口测试,并完成 Init func (that *TestApp) Proj(projects TestProj) *TestApp { that.projs = projects that.ensureInit() return that } // Flows 注册多接口业务流程测试(可选) func (that *TestApp) Flows(flows FlowTest) *TestApp { that.flows = flows return that } // Run 执行测试:注册优雅停测信号 → m.Run → 覆盖率 → Swagger 收尾,返回 exit code func (that *TestApp) Run(m *testing.M) int { that.ensureInit() that.installStopSignals() code := m.Run() if that.Application != nil { _ = that.Db.RollbackTestTx() } that.PrintCoverage() if that.swaggerEnabled { dir := that.swaggerOutputDir() if err := that.GenerateSwagger(that.swaggerTitle, that.swaggerVersion, dir); err != nil { fmt.Printf("Swagger 收尾失败: %v\n", err) } } return code } func (that *TestApp) swaggerOutputDir() string { if that.swaggerDir != "" { return that.swaggerDir } if that.Application != nil && that.Config != nil { if d := that.Config.GetString("tpt"); d != "" { return d } } return "tpt" } func (that *TestApp) installStopSignals() { go func() { sigCh := make(chan os.Signal, 1) signal.Notify(sigCh, os.Interrupt, syscall.SIGTERM) <-sigCh fmt.Println("\n收到停测信号,将在当前接口/流程结束后停止并回滚测试事务…") that.RequestStop() }() } // RequestStop 模拟优雅停测(等同收到 SIGINT):不再开启新的接口/流程,当前跑完后回滚并收尾 func (that *TestApp) RequestStop() { that.stopNewTests.Store(true) } func (that *TestApp) shouldStop() bool { return that.stopNewTests.Load() } func (that *TestApp) ensureInit() { that.initOnce.Do(func() { if that.workDir != "" { if err := os.Chdir(that.workDir); err != nil { fmt.Printf("WorkDir(%s) 失败: %v\n", that.workDir, err) } } if that.configPath == "" { return } origStdout := os.Stdout origStderr := os.Stderr devNull, _ := os.Open(os.DevNull) os.Stdout = devNull os.Stderr = devNull app := Init(that.configPath) os.Stdout = origStdout os.Stderr = origStderr devNull.Close() app.WebConnectLog = nil app.Log.SetOutput(io.Discard) for _, lis := range that.listeners { app.SetConnectListener(lis) } that.Application = app if that.collector == nil { that.collector = &TestCollector{} } if that.projs != nil { router := Router{} for projName, projDef := range that.projs { router[projName] = projDef.Proj } app.SetupForTest(router) } if app.HoTimeCache != nil { app.HoTimeCache.DisableDbCache() } app.Log.SetOutput(os.Stderr) that.inited = true }) } // TestResponse httptest 的原始响应封装 type TestResponse struct { StatusCode int Body []byte Header http.Header } // CoverageReport 覆盖率报告 type CoverageReport struct { Total int Covered int Missing []string Details []MethodCoverage TotalCases int PassedCases int FailedCases int } // MethodCoverage 单个接口的覆盖详情 type MethodCoverage struct { Path string HasTest bool CaseCount int Passed int Failed int Duration time.Duration FailedCases []FailedCase } // FailedCase 失败用例信息 type FailedCase struct { CaseName string FailReason string } // TestRecord 单条测试记录 type TestRecord struct { Path string Desc string CaseName string Passed bool Duration time.Duration Method string Query map[string]interface{} JsonBody interface{} FormBody map[string]interface{} HasFile bool FileField string ResponseBody map[string]interface{} Note string ExpectResult interface{} VerifyError string FailReason string FlowName string FlowGroup string FlowDesc string FlowPrep string FlowStep string FlowStepIndex int BindCase string // FromCase 绑定的单接口验收用例名 } // TestCollector 线程安全的测试记录收集器 type TestCollector struct { mu sync.Mutex Records []TestRecord Visited map[string]bool VisitedFlows map[string]bool } func (c *TestCollector) Add(r TestRecord) { c.mu.Lock() defer c.mu.Unlock() c.Records = append(c.Records, r) } // MarkVisited 标记路径已被测试框架调用 func (c *TestCollector) MarkVisited(path string) { c.mu.Lock() defer c.mu.Unlock() if c.Visited == nil { c.Visited = map[string]bool{} } c.Visited[path] = true } // MarkFlowVisited 标记业务流程已运行 func (c *TestCollector) MarkFlowVisited(name string) { c.mu.Lock() defer c.mu.Unlock() if c.VisitedFlows == nil { c.VisitedFlows = map[string]bool{} } c.VisitedFlows[name] = true } // NewTestApp 立即 Init 的便捷构造(内部委托 Test().Proj()) func NewTestApp(configPath string, projects TestProj, listeners ...func(*Context) bool) *TestApp { return Test(configPath, listeners...).Proj(projects) } // SetupForTest 初始化路由(复用 Run 的路由注册逻辑,但不启动 HTTP 服务) func (that *Application) SetupForTest(router Router) { if that.Router == nil { that.Router = Router{} } for k := range router { v := router[k] if that.Router[k] == nil { that.Router[k] = v } for k1 := range v { v1 := v[k1] if that.Router[k][k1] == nil { that.Router[k][k1] = v1 } for k2 := range v1 { v2 := v1[k2] that.Router[k][k1][k2] = v2 } } } that.MethodRouter = MethodRouter{} modeRouterStrict := true if that.Config.GetBool("modeRouterStrict") == false { modeRouterStrict = false } for pk, pv := range that.Router { if !modeRouterStrict { pk = strings.ToLower(pk) } for ck, cv := range pv { if !modeRouterStrict { ck = strings.ToLower(ck) } for mk, mv := range cv { if !modeRouterStrict { mk = strings.ToLower(mk) } that.MethodRouter["/"+pk+"/"+ck+"/"+mk] = mv } } } } // TestRequest 发送测试 HTTP 请求 func (that *TestApp) TestRequest(method, path string, body *http.Request) TestResponse { w := httptest.NewRecorder() that.Application.ServeHTTP(w, body) return TestResponse{ StatusCode: w.Code, Body: w.Body.Bytes(), Header: w.Header(), } } // RunTests 运行所有注册的测试;方法级事务隔离;另运行 Flows func (that *TestApp) RunTests(t *testing.T) { that.ensureInit() projNames := sortedKeys(that.projs) for _, projName := range projNames { projDef := that.projs[projName] if that.shouldStop() { t.Skip("收到停测信号,跳过剩余测试") return } t.Run(projName, func(t *testing.T) { ctrNames := sortedKeys(projDef.Tests) for _, ctrName := range ctrNames { ctrTest := projDef.Tests[ctrName] if that.shouldStop() { t.Skip("收到停测信号,跳过剩余测试") return } t.Run(ctrName, func(t *testing.T) { methodNames := sortedKeys(ctrTest) for _, methodName := range methodNames { apiTest := ctrTest[methodName] if that.shouldStop() { t.Skip("收到停测信号,跳过剩余接口") return } t.Run(methodName, func(t *testing.T) { path := "/" + projName + "/" + ctrName + "/" + methodName that.collector.MarkVisited(path) err := that.Db.BeginTestTx() if err != nil { t.Fatal("开启测试事务失败:", err) } defer that.Db.RollbackTestTx() methodBuf := &bytes.Buffer{} that.Db.SetTestLogBuffer(methodBuf) defer that.Db.SetTestLogBuffer(nil) api := &Api{ app: that, path: path, t: t, } t.Run(apiTest.Desc, func(t *testing.T) { api.t = t apiTest.Func(api) }) if that.swaggerEnabled { that.flushEndpointSwagger(projName, path) } }) } }) } }) } if len(that.flows) == 0 { return } t.Run("flows", func(t *testing.T) { flowNames := sortedKeys(that.flows) for _, flowName := range flowNames { flowDef := that.flows[flowName] if that.shouldStop() { t.Skip("收到停测信号,跳过剩余流程") return } t.Run(flowName, func(t *testing.T) { that.collector.MarkFlowVisited(flowName) err := that.Db.BeginTestTx() if err != nil { t.Fatal("开启流程测试事务失败:", err) } defer that.Db.RollbackTestTx() methodBuf := &bytes.Buffer{} that.Db.SetTestLogBuffer(methodBuf) defer that.Db.SetTestLogBuffer(nil) flow := &Flow{ app: that, t: t, name: flowName, group: flowDef.Group, desc: flowDef.Desc, prep: flowDef.Prep, } desc := flowDef.Desc if desc == "" { desc = flowName } t.Run(desc, func(t *testing.T) { flow.t = t flowDef.Func(flow) }) if that.swaggerEnabled { that.flushFlowSwagger(flowName) } }) } }) } func sortedKeys[M ~map[string]V, V any](m M) []string { keys := make([]string, 0, len(m)) for k := range m { keys = append(keys, k) } sort.Strings(keys) return keys } // PrintCoverage 输出 API 测试覆盖率报告 func (that *TestApp) PrintCoverage() CoverageReport { report := CoverageReport{} pathRecords := map[string][]TestRecord{} that.collector.mu.Lock() for _, r := range that.collector.Records { if r.FlowName != "" { continue } pathRecords[r.Path] = append(pathRecords[r.Path], r) } visited := that.collector.Visited that.collector.mu.Unlock() for projName, projDef := range that.projs { for ctrName, ctr := range projDef.Proj { for methodName := range ctr { path := "/" + projName + "/" + ctrName + "/" + methodName report.Total++ ctrTest, hasCtr := projDef.Tests[ctrName] if !hasCtr { report.Missing = append(report.Missing, projName+"/"+ctrName+"/"+methodName) report.Details = append(report.Details, MethodCoverage{Path: path, HasTest: false}) continue } _, hasMethod := ctrTest[methodName] if !hasMethod { report.Missing = append(report.Missing, projName+"/"+ctrName+"/"+methodName) report.Details = append(report.Details, MethodCoverage{Path: path, HasTest: false}) continue } report.Covered++ mc := MethodCoverage{Path: path, HasTest: true} if records, ok := pathRecords[path]; ok { mc.CaseCount = len(records) for _, r := range records { mc.Duration += r.Duration if r.Passed { mc.Passed++ } else { mc.Failed++ mc.FailedCases = append(mc.FailedCases, FailedCase{ CaseName: r.CaseName, FailReason: r.FailReason, }) } } } report.Details = append(report.Details, mc) } } } for _, d := range report.Details { report.TotalCases += d.CaseCount report.PassedCases += d.Passed report.FailedCases += d.Failed } ranCount := len(visited) isPartialRun := ranCount > 0 && ranCount < report.Covered var passedDetails, failedDetails []MethodCoverage for _, d := range report.Details { if !d.HasTest || d.CaseCount == 0 { continue } if isPartialRun && !visited[d.Path] { continue } if d.Failed > 0 { failedDetails = append(failedDetails, d) } else { passedDetails = append(passedDetails, d) } } ranInterfaces := len(failedDetails) + len(passedDetails) var ranTotalCases, ranPassedCases, ranFailedCases int if isPartialRun { for _, d := range failedDetails { ranTotalCases += d.CaseCount ranPassedCases += d.Passed ranFailedCases += d.Failed } for _, d := range passedDetails { ranTotalCases += d.CaseCount ranPassedCases += d.Passed ranFailedCases += d.Failed } } else { ranTotalCases = report.TotalCases ranPassedCases = report.PassedCases ranFailedCases = report.FailedCases ranInterfaces = len(failedDetails) + len(passedDetails) } fmt.Println("\n========== API 测试覆盖率报告 ==========") if report.Total > 0 { fmt.Printf("总接口: %d | 已覆盖: %d | 覆盖率: %.1f%%\n", report.Total, report.Covered, float64(report.Covered)/float64(report.Total)*100) } else { fmt.Println("总接口: 0") } if isPartialRun && ranTotalCases > 0 { passRate := float64(ranPassedCases) / float64(ranTotalCases) * 100 fmt.Printf("本次运行: %d 个接口 | 总用例: %d | 通过: %d | 未通过: %d | 通过率: %.1f%%\n", ranInterfaces, ranTotalCases, ranPassedCases, ranFailedCases, passRate) } else if ranTotalCases > 0 { passRate := float64(ranPassedCases) / float64(ranTotalCases) * 100 fmt.Printf("总用例: %d | 通过: %d | 未通过: %d | 通过率: %.1f%%\n", ranTotalCases, ranPassedCases, ranFailedCases, passRate) } if len(failedDetails) > 0 { fmt.Printf("\n---------- 未通过的接口 (%d个) ----------\n", len(failedDetails)) for _, d := range failedDetails { fmt.Printf(" ✗ %s\t%d 用例 (%d 通过, %d 失败) %v\n", d.Path, d.CaseCount, d.Passed, d.Failed, d.Duration) for _, fc := range d.FailedCases { if fc.FailReason != "" { fmt.Printf(" [失败] %s — %s\n", fc.CaseName, fc.FailReason) } else { fmt.Printf(" [失败] %s\n", fc.CaseName) } } } } if len(passedDetails) > 0 { fmt.Println("\n---------- 通过的接口 ----------") for _, d := range passedDetails { fmt.Printf(" ✓ %s\t%d 用例 (%d 通过, %d 失败) %v\n", d.Path, d.CaseCount, d.Passed, d.Failed, d.Duration) } } if len(report.Missing) > 0 { fmt.Printf("\n未覆盖的接口: (%d 个)\n", len(report.Missing)) for _, path := range report.Missing { fmt.Println(" -", path) } } fmt.Println("========================================") return report } // DB 获取测试应用的数据库实例(带 testTx) func (that *TestApp) DB() *HoTimeDB { return &that.Db }