Source file src/net/http/export_test.go

     1  // Copyright 2011 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     4  
     5  // Bridge package to expose http internals to tests in the http_test
     6  // package.
     7  
     8  package http
     9  
    10  import (
    11  	"context"
    12  	"fmt"
    13  	"net"
    14  	"net/url"
    15  	"slices"
    16  	"sync"
    17  	"testing"
    18  	"time"
    19  )
    20  
    21  var (
    22  	DefaultUserAgent                  = defaultUserAgent
    23  	NewLoggingConn                    = newLoggingConn
    24  	ExportRefererForURL               = refererForURL
    25  	ExportServerNewConn               = (*Server).newConn
    26  	ExportCloseWriteAndWait           = (*conn).closeWriteAndWait
    27  	ExportErrRequestCanceled          = errRequestCanceled
    28  	ExportErrRequestCanceledConn      = errRequestCanceledConn
    29  	ExportErrServerClosedIdle         = errServerClosedIdle
    30  	ExportServeFile                   = serveFile
    31  	ExportScanETag                    = scanETag
    32  	Export_shouldCopyHeaderOnRedirect = shouldCopyHeaderOnRedirect
    33  	Export_writeStatusLine            = writeStatusLine
    34  	Export_is408Message               = is408Message
    35  	MaxPostCloseReadTime              = maxPostCloseReadTime
    36  	ProtocolSetHTTP3                  = protocolSetHTTP3
    37  )
    38  
    39  var MaxWriteWaitBeforeConnReuse = &maxWriteWaitBeforeConnReuse
    40  
    41  func init() {
    42  	// We only want to pay for this cost during testing.
    43  	// When not under test, these values are always nil
    44  	// and never assigned to.
    45  	testHookMu = new(sync.Mutex)
    46  
    47  	testHookClientDoResult = func(res *Response, err error) {
    48  		if err != nil {
    49  			if _, ok := err.(*url.Error); !ok {
    50  				panic(fmt.Sprintf("unexpected Client.Do error of type %T; want *url.Error", err))
    51  			}
    52  		} else {
    53  			if res == nil {
    54  				panic("Client.Do returned nil, nil")
    55  			}
    56  			if res.Body == nil {
    57  				panic("Client.Do returned nil res.Body and no error")
    58  			}
    59  		}
    60  	}
    61  }
    62  
    63  func CondSkipHTTP2(t testing.TB) {
    64  	if omitBundledHTTP2 {
    65  		t.Skip("skipping HTTP/2 test when nethttpomithttp2 build tag in use")
    66  	}
    67  }
    68  
    69  var (
    70  	SetEnterRoundTripHook = hookSetter(&testHookEnterRoundTrip)
    71  	SetRoundTripRetried   = hookSetter(&testHookRoundTripRetried)
    72  )
    73  
    74  func SetReadLoopBeforeNextReadHook(f func()) {
    75  	unnilTestHook(&f)
    76  	testHookReadLoopBeforeNextRead = f
    77  }
    78  
    79  // SetPendingDialHooks sets the hooks that run before and after handling
    80  // pending dials.
    81  func SetPendingDialHooks(before, after func()) {
    82  	unnilTestHook(&before)
    83  	unnilTestHook(&after)
    84  	testHookPrePendingDial, testHookPostPendingDial = before, after
    85  }
    86  
    87  func SetTestHookServerServe(fn func(*Server, net.Listener)) { testHookServerServe = fn }
    88  
    89  func SetTestHookProxyConnectTimeout(t *testing.T, f func(context.Context, time.Duration) (context.Context, context.CancelFunc)) {
    90  	orig := testHookProxyConnectTimeout
    91  	t.Cleanup(func() {
    92  		testHookProxyConnectTimeout = orig
    93  	})
    94  	testHookProxyConnectTimeout = f
    95  }
    96  
    97  func NewTestTimeoutHandler(handler Handler, ctx context.Context) Handler {
    98  	return &timeoutHandler{
    99  		handler:     handler,
   100  		testContext: ctx,
   101  		// (no body)
   102  	}
   103  }
   104  
   105  func ResetCachedEnvironment() {
   106  	resetProxyConfig()
   107  }
   108  
   109  func (t *Transport) NumPendingRequestsForTesting() int {
   110  	t.reqMu.Lock()
   111  	defer t.reqMu.Unlock()
   112  	return len(t.reqCanceler)
   113  }
   114  
   115  func (t *Transport) IdleConnKeysForTesting() (keys []string) {
   116  	keys = make([]string, 0)
   117  	t.idleMu.Lock()
   118  	defer t.idleMu.Unlock()
   119  	for key := range t.idleConn {
   120  		keys = append(keys, key.String())
   121  	}
   122  	slices.Sort(keys)
   123  	return
   124  }
   125  
   126  func (t *Transport) IdleConnKeyCountForTesting() int {
   127  	t.idleMu.Lock()
   128  	defer t.idleMu.Unlock()
   129  	return len(t.idleConn)
   130  }
   131  
   132  func (t *Transport) IdleConnStrsForTesting() []string {
   133  	var ret []string
   134  	t.idleMu.Lock()
   135  	defer t.idleMu.Unlock()
   136  	for _, conns := range t.idleConn {
   137  		for _, pc := range conns {
   138  			if pc.conn == nil {
   139  				continue
   140  			}
   141  			ret = append(ret, pc.conn.LocalAddr().String()+"/"+pc.conn.RemoteAddr().String())
   142  		}
   143  	}
   144  	if t.h2Transport != nil {
   145  		ret = append(ret, t.h2Transport.IdleConnStrsForTesting()...)
   146  	}
   147  	slices.Sort(ret)
   148  	return ret
   149  }
   150  
   151  func (t *Transport) IdleConnCountForTesting(scheme, addr string) int {
   152  	t.idleMu.Lock()
   153  	defer t.idleMu.Unlock()
   154  	key := connectMethodKey{"", scheme, addr, false}
   155  	cacheKey := key.String()
   156  	for k, conns := range t.idleConn {
   157  		if k.String() == cacheKey {
   158  			return len(conns)
   159  		}
   160  	}
   161  	return 0
   162  }
   163  
   164  func (t *Transport) IdleConnWaitMapSizeForTesting() int {
   165  	t.idleMu.Lock()
   166  	defer t.idleMu.Unlock()
   167  	return len(t.idleConnWait)
   168  }
   169  
   170  func (t *Transport) IsIdleForTesting() bool {
   171  	t.idleMu.Lock()
   172  	defer t.idleMu.Unlock()
   173  	return t.closeIdle
   174  }
   175  
   176  func (t *Transport) QueueForIdleConnForTesting() {
   177  	t.queueForIdleConn(nil)
   178  }
   179  
   180  // PutIdleTestConn reports whether it was able to insert a fresh
   181  // persistConn for scheme, addr into the idle connection pool.
   182  func (t *Transport) PutIdleTestConn(scheme, addr string) bool {
   183  	c, _ := net.Pipe()
   184  	key := connectMethodKey{"", scheme, addr, false}
   185  
   186  	if t.MaxConnsPerHost > 0 {
   187  		// Transport is tracking conns-per-host.
   188  		// Increment connection count to account
   189  		// for new persistConn created below.
   190  		t.connsPerHostMu.Lock()
   191  		if t.connsPerHost == nil {
   192  			t.connsPerHost = make(map[connectMethodKey]int)
   193  		}
   194  		t.connsPerHost[key]++
   195  		t.connsPerHostMu.Unlock()
   196  	}
   197  
   198  	return t.tryPutIdleConn(&persistConn{
   199  		t:        t,
   200  		conn:     c,                   // dummy
   201  		closech:  make(chan struct{}), // so it can be closed
   202  		cacheKey: key,
   203  	}) == nil
   204  }
   205  
   206  // PutIdleTestConnH2 reports whether it was able to insert a fresh
   207  // HTTP/2 persistConn for scheme, addr into the idle connection pool.
   208  func (t *Transport) PutIdleTestConnH2(scheme, addr string, alt RoundTripper) bool {
   209  	key := connectMethodKey{"", scheme, addr, false}
   210  
   211  	if t.MaxConnsPerHost > 0 {
   212  		// Transport is tracking conns-per-host.
   213  		// Increment connection count to account
   214  		// for new persistConn created below.
   215  		t.connsPerHostMu.Lock()
   216  		if t.connsPerHost == nil {
   217  			t.connsPerHost = make(map[connectMethodKey]int)
   218  		}
   219  		t.connsPerHost[key]++
   220  		t.connsPerHostMu.Unlock()
   221  	}
   222  
   223  	return t.tryPutIdleConn(&persistConn{
   224  		t:        t,
   225  		alt:      alt,
   226  		cacheKey: key,
   227  	}) == nil
   228  }
   229  
   230  // All test hooks must be non-nil so they can be called directly,
   231  // but the tests use nil to mean hook disabled.
   232  func unnilTestHook(f *func()) {
   233  	if *f == nil {
   234  		*f = nop
   235  	}
   236  }
   237  
   238  func hookSetter(dst *func()) func(func()) {
   239  	return func(fn func()) {
   240  		unnilTestHook(&fn)
   241  		*dst = fn
   242  	}
   243  }
   244  
   245  func (s *Server) ExportAllConnsIdle() bool {
   246  	s.mu.Lock()
   247  	defer s.mu.Unlock()
   248  	for c := range s.activeConn {
   249  		st, unixSec := c.getState()
   250  		if unixSec == 0 || st != StateIdle {
   251  			return false
   252  		}
   253  	}
   254  	return true
   255  }
   256  
   257  func (s *Server) ExportAllConnsByState() map[ConnState]int {
   258  	states := map[ConnState]int{}
   259  	s.mu.Lock()
   260  	defer s.mu.Unlock()
   261  	for c := range s.activeConn {
   262  		st, _ := c.getState()
   263  		states[st] += 1
   264  	}
   265  	return states
   266  }
   267  
   268  func (r *Request) WithT(t *testing.T) *Request {
   269  	return r.WithContext(context.WithValue(r.Context(), tLogKey{}, t.Logf))
   270  }
   271  
   272  func (r *Request) ExportIsReplayable() bool { return r.isReplayable() }
   273  
   274  // ExportCloseTransportConnsAbruptly closes all idle connections from
   275  // tr in an abrupt way, just reaching into the underlying Conns and
   276  // closing them, without telling the Transport or its persistConns
   277  // that it's doing so. This is to simulate the server closing connections
   278  // on the Transport.
   279  func ExportCloseTransportConnsAbruptly(tr *Transport) {
   280  	tr.idleMu.Lock()
   281  	for _, pcs := range tr.idleConn {
   282  		for _, pc := range pcs {
   283  			pc.conn.Close()
   284  		}
   285  	}
   286  	tr.idleMu.Unlock()
   287  }
   288  
   289  // ResponseWriterConnForTesting returns w's underlying connection, if w
   290  // is a regular *response ResponseWriter.
   291  func ResponseWriterConnForTesting(w ResponseWriter) (c net.Conn, ok bool) {
   292  	if r, ok := w.(*response); ok {
   293  		return r.conn.rwc, true
   294  	}
   295  	return nil, false
   296  }
   297  
   298  func init() {
   299  	// Set the default rstAvoidanceDelay to the minimum possible value to shake
   300  	// out tests that unexpectedly depend on it. Such tests should use
   301  	// runTimeSensitiveTest and SetRSTAvoidanceDelay to explicitly raise the delay
   302  	// if needed.
   303  	rstAvoidanceDelay = 1 * time.Nanosecond
   304  }
   305  
   306  // SetRSTAvoidanceDelay sets how long we are willing to wait between calling
   307  // CloseWrite on a connection and fully closing the connection.
   308  func SetRSTAvoidanceDelay(t *testing.T, d time.Duration) {
   309  	prevDelay := rstAvoidanceDelay
   310  	t.Cleanup(func() {
   311  		rstAvoidanceDelay = prevDelay
   312  	})
   313  	rstAvoidanceDelay = d
   314  }
   315  

View as plain text