package middleware import ( "bytes" "io" "testing" "time" ) // TestRateLimitedReaderAtRate 验证长读大致符合设定速率。 // 设 100 KB/s,读取 400 KB 总字节,期望耗时约 4s(±1s 容忍)。 func TestRateLimitedReaderAtRate(t *testing.T) { const rate = 100 * 1024 const total = 400 * 1024 src := io.LimitReader(bytes.NewReader(make([]byte, total+1024)), total) rl := NewRateLimitedReader(src, rate).(*rateLimitedReader) buf := make([]byte, 32*1024) // 32KB 块 start := time.Now() read := 0 for { n, err := rl.Read(buf) read += n if err == io.EOF { break } if err != nil { t.Fatalf("unexpected err: %v", err) } } elapsed := time.Since(start) want := time.Duration(float64(time.Second) * float64(total) / float64(rate)) if elapsed < want-time.Second { t.Fatalf("读取过快 elapsed=%v want>=%v", elapsed, want) } if elapsed > want+1500*time.Millisecond { t.Fatalf("读取过慢 elapsed=%v want<=%v", elapsed, want+1500*time.Millisecond) } if read != total { t.Fatalf("读到 %d 字节,期望 %d", read, total) } } // TestRateLimitedReaderZeroPassthrough 速率 0 时不应引入任何延迟/包封。 func TestRateLimitedReaderZeroPassthrough(t *testing.T) { src := bytes.NewReader([]byte("hello")) rl := NewRateLimitedReader(src, 0) if rl == src { // 透传:返回原 reader } else if _, ok := rl.(*rateLimitedReader); ok { // 0 速率时按实现可走 enabled=false(不退化亦可) } // 关键:必须能读完 b, err := io.ReadAll(rl) if err != nil || string(b) != "hello" { t.Fatalf("0 速率透传失败: %q %v", b, err) } } // TestWrapReadCloserClose 验证包裹后 Close 透传到底层。 func TestWrapReadCloserClose(t *testing.T) { src := &closeCount{Reader: bytes.NewReader([]byte("xyz")), closed: 0} rc := WrapReadCloser(src, 50*1024) if rc == nil { t.Fatal("WrapReadCloser nil") } _, _ = io.ReadAll(rc) if err := rc.Close(); err != nil { t.Fatalf("close err: %v", err) } if src.closed != 1 { t.Fatalf("底层 Close 未被调用: %d", src.closed) } } type closeCount struct { io.Reader closed int } func (c *closeCount) Close() error { c.closed++; return nil }