package hotime import ( "bytes" "fmt" "io" "net/http" "net/http/httptest" "os" "strings" "sync" "testing" "time" . "code.hoteas.com/golang/hotime/db" ) // TestApp 测试应用,封装 Application 提供测试能力 type TestApp struct { *Application projs TestProj collector *TestCollector } // 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 } // TestCollector 线程安全的测试记录收集器 type TestCollector struct { mu sync.Mutex Records []TestRecord Visited map[string]bool // 被 RunTests 实际调用过的路径(不论是否产生记录) } func (c *TestCollector) Add(r TestRecord) { c.mu.Lock() defer c.mu.Unlock() c.Records = append(c.Records, r) } // MarkVisited 标记路径已被测试框架调用(由 RunTests 在进入方法级 t.Run 时调用) 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 } // NewTestApp 创建测试应用实例 // configPath: 配置文件路径(如 "../config/config.json") // projects: 项目定义(路由 + 测试) // listeners: 可选的请求拦截器(与 SetConnectListener 相同) func NewTestApp(configPath string, projects TestProj, listeners ...func(*Context) bool) *TestApp { origStdout := os.Stdout origStderr := os.Stderr devNull, _ := os.Open(os.DevNull) os.Stdout = devNull os.Stderr = devNull app := Init(configPath) os.Stdout = origStdout os.Stderr = origStderr devNull.Close() app.WebConnectLog = nil app.Log.SetOutput(io.Discard) for _, lis := range listeners { app.SetConnectListener(lis) } router := Router{} for projName, projDef := range projects { router[projName] = projDef.Proj } app.SetupForTest(router) if app.HoTimeCache != nil { app.HoTimeCache.DisableDbCache() } app.Log.SetOutput(os.Stderr) return &TestApp{ Application: app, projs: projects, collector: &TestCollector{}, } } // 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 运行所有注册的测试,每个方法级别使用事务隔离 func (that *TestApp) RunTests(t *testing.T) { for projName, projDef := range that.projs { projName := projName projDef := projDef t.Run(projName, func(t *testing.T) { for ctrName, ctrTest := range projDef.Tests { ctrName := ctrName ctrTest := ctrTest t.Run(ctrName, func(t *testing.T) { for methodName, apiTest := range ctrTest { methodName := methodName apiTest := apiTest 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) }) }) } }) } }) } } // PrintCoverage 输出 API 测试覆盖率报告 func (that *TestApp) PrintCoverage() CoverageReport { report := CoverageReport{} pathRecords := map[string][]TestRecord{} that.collector.mu.Lock() for _, r := range that.collector.Records { 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 }