Files
hotime/testing_helper.go
T

403 lines
10 KiB
Go
Raw Normal View History

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
}