Files
hotime/testing_helper.go
T
hoteas 8a1a27b457 feat(testing): Listener 链式入口并精简测试文档运行指引
将 connect listener 并入 Test 链、删除 NewTestApp,正文只示范 -run 单测,全量命令仅保留在文末 CI 节。

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-07-13 19:56:36 +08:00

632 lines
15 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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) *TestApp {
return &TestApp{
configPath: configPath,
collector: &TestCollector{},
}
}
// WorkDir 设置工作目录(在 Init 前 chdir)
func (that *TestApp) WorkDir(dir string) *TestApp {
that.workDir = dir
return that
}
// Listener 追加连接监听(可选;可多次调用或一次传多个)。
// Init 前仅暂存,ensureInit 时批量 SetConnectListenerInit 后立刻挂上。
// 返回 true 表示已处理并停止进入控制器,false 继续。
func (that *TestApp) Listener(fns ...func(*Context) bool) *TestApp {
that.listeners = append(that.listeners, fns...)
if that.inited && that.Application != nil {
for _, fn := range fns {
that.SetConnectListener(fn)
}
}
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
}
// 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
}