1
2
3
4
5 package http2_test
6
7 import (
8 "bytes"
9 "compress/gzip"
10 "compress/zlib"
11 "context"
12 "crypto/tls"
13 "crypto/x509"
14 "errors"
15 "flag"
16 "fmt"
17 "io"
18 "log"
19 "math"
20 "net"
21 "net/http"
22 "net/http/httptest"
23 "os"
24 "reflect"
25 "runtime"
26 "slices"
27 "strconv"
28 "strings"
29 "sync"
30 "testing"
31 "testing/synctest"
32 "time"
33 _ "unsafe"
34
35 "net/http/internal/http2"
36 . "net/http/internal/http2"
37 "net/http/internal/testcert"
38
39 "golang.org/x/net/http2/hpack"
40 )
41
42 var stderrVerbose = flag.Bool("stderr_verbose", false, "Mirror verbosity to stderr, unbuffered")
43
44 func stderrv() io.Writer {
45 if *stderrVerbose {
46 return os.Stderr
47 }
48
49 return io.Discard
50 }
51
52 type safeBuffer struct {
53 b bytes.Buffer
54 m sync.Mutex
55 }
56
57 func (sb *safeBuffer) Write(d []byte) (int, error) {
58 sb.m.Lock()
59 defer sb.m.Unlock()
60 return sb.b.Write(d)
61 }
62
63 func (sb *safeBuffer) Bytes() []byte {
64 sb.m.Lock()
65 defer sb.m.Unlock()
66 return sb.b.Bytes()
67 }
68
69 func (sb *safeBuffer) Len() int {
70 sb.m.Lock()
71 defer sb.m.Unlock()
72 return sb.b.Len()
73 }
74
75 type serverTester struct {
76 cc net.Conn
77 t *testing.T
78 h1server *http.Server
79 h2server *Server
80 serverLogBuf safeBuffer
81 logFilter []string
82 scMu sync.Mutex
83 sc *ServerConn
84 wrotePreface bool
85 testConnFramer
86
87 callsMu sync.Mutex
88 calls []*serverHandlerCall
89
90
91
92
93
94 frameReadLogMu sync.Mutex
95 frameReadLogBuf bytes.Buffer
96 frameWriteLogMu sync.Mutex
97 frameWriteLogBuf bytes.Buffer
98
99
100 headerBuf bytes.Buffer
101 hpackEnc *hpack.Encoder
102 }
103
104 type twriter struct {
105 t testing.TB
106 st *serverTester
107 }
108
109 func (w twriter) Write(p []byte) (n int, err error) {
110 if w.st != nil {
111 ps := string(p)
112 for _, phrase := range w.st.logFilter {
113 if strings.Contains(ps, phrase) {
114 return len(p), nil
115 }
116 }
117 }
118 w.t.Logf("%s", p)
119 return len(p), nil
120 }
121
122 func newTestServer(t testing.TB, handler http.HandlerFunc, opts ...any) *httptest.Server {
123 t.Helper()
124 if handler == nil {
125 handler = func(w http.ResponseWriter, req *http.Request) {}
126 }
127 ts := httptest.NewUnstartedServer(handler)
128 ts.EnableHTTP2 = true
129 ts.Config.ErrorLog = log.New(twriter{t: t}, "", log.LstdFlags)
130 ts.Config.Protocols = protocols("h2")
131 for _, opt := range opts {
132 switch v := opt.(type) {
133 case func(*httptest.Server):
134 v(ts)
135 case func(*http.Server):
136 v(ts.Config)
137 case func(*http.HTTP2Config):
138 if ts.Config.HTTP2 == nil {
139 ts.Config.HTTP2 = &http.HTTP2Config{}
140 }
141 v(ts.Config.HTTP2)
142 default:
143 t.Fatalf("unknown newTestServer option type %T", v)
144 }
145 }
146
147 if ts.Config.Protocols.HTTP2() {
148 ts.TLS = testServerTLSConfig
149 if ts.Config.TLSConfig != nil {
150 ts.TLS = ts.Config.TLSConfig
151 }
152 ts.StartTLS()
153 } else if ts.Config.Protocols.UnencryptedHTTP2() {
154 ts.EnableHTTP2 = false
155 ts.Start()
156 } else {
157 t.Fatalf("Protocols contains neither HTTP2 nor UnencryptedHTTP2")
158 }
159
160 t.Cleanup(func() {
161 ts.CloseClientConnections()
162 ts.Close()
163 })
164
165 return ts
166 }
167
168 type serverTesterOpt string
169
170 var optFramerReuseFrames = serverTesterOpt("frame_reuse_frames")
171
172 var optQuiet = func(server *http.Server) {
173 server.ErrorLog = log.New(io.Discard, "", 0)
174 }
175
176 func newServerTester(t *testing.T, handler http.HandlerFunc, opts ...any) *serverTester {
177 t.Helper()
178
179 h1server := &http.Server{}
180 var tlsState *tls.ConnectionState
181 for _, opt := range opts {
182 switch v := opt.(type) {
183 case func(*http.Server):
184 v(h1server)
185 case func(*http.HTTP2Config):
186 if h1server.HTTP2 == nil {
187 h1server.HTTP2 = &http.HTTP2Config{}
188 }
189 v(h1server.HTTP2)
190 case func(*tls.ConnectionState):
191 if tlsState == nil {
192 tlsState = &tls.ConnectionState{
193 Version: tls.VersionTLS13,
194 ServerName: "go.dev",
195 CipherSuite: tls.TLS_AES_128_GCM_SHA256,
196 }
197 }
198 v(tlsState)
199 default:
200 t.Fatalf("unknown newServerTester option type %T", v)
201 }
202 }
203
204 tlsConfig := h1server.TLSConfig
205 if tlsConfig == nil {
206 cert, err := tls.X509KeyPair(testcert.LocalhostCert, testcert.LocalhostKey)
207 if err != nil {
208 t.Fatal(err)
209 }
210 tlsConfig = &tls.Config{
211 Certificates: []tls.Certificate{cert},
212 InsecureSkipVerify: true,
213 NextProtos: []string{"h2"},
214 }
215 h1server.TLSConfig = tlsConfig
216 }
217
218 var cli, srv net.Conn
219
220 cliPipe, srvPipe := synctestNetPipe()
221
222 if h1server.Protocols != nil && h1server.Protocols.UnencryptedHTTP2() {
223 cli, srv = cliPipe, srvPipe
224 } else {
225 cli = tls.Client(cliPipe, &tls.Config{
226 InsecureSkipVerify: true,
227 NextProtos: []string{"h2"},
228 })
229 srv = tls.Server(srvPipe, tlsConfig)
230 }
231
232 st := &serverTester{
233 t: t,
234 cc: cli,
235 h1server: h1server,
236 }
237 st.hpackEnc = hpack.NewEncoder(&st.headerBuf)
238 if h1server.ErrorLog == nil {
239 h1server.ErrorLog = log.New(io.MultiWriter(stderrv(), twriter{t: t, st: st}, &st.serverLogBuf), "", log.LstdFlags)
240 }
241
242 if handler == nil {
243 handler = serverTesterHandler{st}.ServeHTTP
244 }
245 h1server.Handler = handler
246
247 t.Cleanup(func() {
248 st.Close()
249 time.Sleep(GoAwayTimeout)
250 })
251
252 connc := make(chan *ServerConn)
253 h1server.ConnContext = func(ctx context.Context, conn net.Conn) context.Context {
254 ctx = context.WithValue(ctx, NewConnContextKey, func(sc *ServerConn) {
255 connc <- sc
256 })
257 if tlsState != nil {
258 ctx = context.WithValue(ctx, ConnectionStateContextKey, func() tls.ConnectionState {
259 return *tlsState
260 })
261 }
262 return ctx
263 }
264 go func() {
265 li := newOneConnListener(srv)
266 t.Cleanup(func() {
267 li.Close()
268 })
269 h1server.Serve(li)
270 }()
271 if cliTLS, ok := cli.(*tls.Conn); ok {
272 if err := cliTLS.Handshake(); err != nil {
273 t.Fatalf("client TLS handshake: %v", err)
274 }
275 cliTLS.SetReadDeadline(time.Now())
276 } else {
277
278
279 st.writePreface()
280 st.wrotePreface = true
281 cliPipe.SetReadDeadline(time.Now())
282 }
283 st.sc = <-connc
284
285 st.fr = NewFramer(st.cc, st.cc)
286 st.testConnFramer = testConnFramer{
287 t: t,
288 fr: NewFramer(cli, cli),
289 dec: hpack.NewDecoder(InitialHeaderTableSize, nil),
290 }
291 synctest.Wait()
292 return st
293 }
294
295 type netConnWithConnectionState struct {
296 net.Conn
297 state tls.ConnectionState
298 }
299
300 func (c *netConnWithConnectionState) ConnectionState() tls.ConnectionState {
301 return c.state
302 }
303
304 func (c *netConnWithConnectionState) HandshakeContext() tls.ConnectionState {
305 return c.state
306 }
307
308 type serverTesterHandler struct {
309 st *serverTester
310 }
311
312 func (h serverTesterHandler) ServeHTTP(w http.ResponseWriter, req *http.Request) {
313 call := &serverHandlerCall{
314 w: w,
315 req: req,
316 ch: make(chan func()),
317 }
318 h.st.t.Cleanup(call.exit)
319 h.st.callsMu.Lock()
320 h.st.calls = append(h.st.calls, call)
321 h.st.callsMu.Unlock()
322 for f := range call.ch {
323 f()
324 }
325 }
326
327
328 type serverHandlerCall struct {
329 w http.ResponseWriter
330 req *http.Request
331 closeOnce sync.Once
332 ch chan func()
333 }
334
335
336 func (call *serverHandlerCall) do(f func(http.ResponseWriter, *http.Request)) {
337 donec := make(chan struct{})
338 call.ch <- func() {
339 defer close(donec)
340 f(call.w, call.req)
341 }
342 <-donec
343 }
344
345
346 func (call *serverHandlerCall) exit() {
347 call.closeOnce.Do(func() {
348 close(call.ch)
349 })
350 }
351
352
353 func (st *serverTester) sync() {
354 synctest.Wait()
355 }
356
357
358 func (st *serverTester) advance(d time.Duration) {
359 time.Sleep(d)
360 synctest.Wait()
361 }
362
363 func (st *serverTester) authority() string {
364 return "dummy.tld"
365 }
366
367 func (st *serverTester) addLogFilter(phrase string) {
368 st.logFilter = append(st.logFilter, phrase)
369 }
370
371 func (st *serverTester) nextHandlerCall() *serverHandlerCall {
372 st.t.Helper()
373 synctest.Wait()
374 st.callsMu.Lock()
375 defer st.callsMu.Unlock()
376 if len(st.calls) == 0 {
377 st.t.Fatal("expected server handler call, got none")
378 }
379 call := st.calls[0]
380 st.calls = st.calls[1:]
381 return call
382 }
383
384 func (st *serverTester) streamExists(id uint32) bool {
385 return st.sc.TestStreamExists(id)
386 }
387
388 func (st *serverTester) streamState(id uint32) StreamState {
389 return st.sc.TestStreamState(id)
390 }
391
392 func (st *serverTester) Close() {
393 if st.t.Failed() {
394 st.frameReadLogMu.Lock()
395 if st.frameReadLogBuf.Len() > 0 {
396 st.t.Logf("Framer read log:\n%s", st.frameReadLogBuf.String())
397 }
398 st.frameReadLogMu.Unlock()
399
400 st.frameWriteLogMu.Lock()
401 if st.frameWriteLogBuf.Len() > 0 {
402 st.t.Logf("Framer write log:\n%s", st.frameWriteLogBuf.String())
403 }
404 st.frameWriteLogMu.Unlock()
405
406
407
408
409
410 if st.cc != nil {
411 st.cc.Close()
412 }
413 }
414 if st.cc != nil {
415 st.cc.Close()
416 }
417 log.SetOutput(os.Stderr)
418 }
419
420
421
422 func (st *serverTester) greet() {
423 st.t.Helper()
424 st.greetAndCheckSettings(func(Setting) error { return nil })
425 }
426
427 func (st *serverTester) greetAndCheckSettings(checkSetting func(s Setting) error) {
428 st.t.Helper()
429 st.writePreface()
430 st.writeSettings()
431 st.sync()
432 readFrame[*SettingsFrame](st.t, st).ForeachSetting(checkSetting)
433 st.writeSettingsAck()
434
435
436 var gotSettingsAck bool
437 var gotWindowUpdate bool
438
439 for range 2 {
440 f := st.readFrame()
441 if f == nil {
442 st.t.Fatal("wanted a settings ACK and window update, got none")
443 }
444 switch f := f.(type) {
445 case *SettingsFrame:
446 if !f.Header().Flags.Has(FlagSettingsAck) {
447 st.t.Fatal("Settings Frame didn't have ACK set")
448 }
449 gotSettingsAck = true
450
451 case *WindowUpdateFrame:
452 if f.FrameHeader.StreamID != 0 {
453 st.t.Fatalf("WindowUpdate StreamID = %d; want 0", f.FrameHeader.StreamID)
454 }
455 gotWindowUpdate = true
456
457 default:
458 st.t.Fatalf("Wanting a settings ACK or window update, received a %T", f)
459 }
460 }
461
462 if !gotSettingsAck {
463 st.t.Fatalf("Didn't get a settings ACK")
464 }
465 if !gotWindowUpdate {
466 st.t.Fatalf("Didn't get a window update")
467 }
468 }
469
470 func (st *serverTester) writePreface() {
471 if st.wrotePreface {
472 return
473 }
474 n, err := st.cc.Write([]byte(ClientPreface))
475 if err != nil {
476 st.t.Fatalf("Error writing client preface: %v", err)
477 }
478 if n != len(ClientPreface) {
479 st.t.Fatalf("Writing client preface, wrote %d bytes; want %d", n, len(ClientPreface))
480 }
481 }
482
483 func (st *serverTester) encodeHeaderField(k, v string) {
484 err := st.hpackEnc.WriteField(hpack.HeaderField{Name: k, Value: v})
485 if err != nil {
486 st.t.Fatalf("HPACK encoding error for %q/%q: %v", k, v, err)
487 }
488 }
489
490
491
492 func (st *serverTester) encodeHeaderRaw(headers ...string) []byte {
493 if len(headers)%2 == 1 {
494 panic("odd number of kv args")
495 }
496 st.headerBuf.Reset()
497 for len(headers) > 0 {
498 k, v := headers[0], headers[1]
499 st.encodeHeaderField(k, v)
500 headers = headers[2:]
501 }
502 return st.headerBuf.Bytes()
503 }
504
505
506
507
508
509
510 func (st *serverTester) encodeHeader(headers ...string) []byte {
511 if len(headers)%2 == 1 {
512 panic("odd number of kv args")
513 }
514
515 st.headerBuf.Reset()
516 defaultAuthority := st.authority()
517
518 if len(headers) == 0 {
519
520
521 st.encodeHeaderField(":method", "GET")
522 st.encodeHeaderField(":scheme", "https")
523 st.encodeHeaderField(":authority", defaultAuthority)
524 st.encodeHeaderField(":path", "/")
525 return st.headerBuf.Bytes()
526 }
527
528 if len(headers) == 2 && headers[0] == ":method" {
529
530 st.encodeHeaderField(":method", headers[1])
531 st.encodeHeaderField(":scheme", "https")
532 st.encodeHeaderField(":authority", defaultAuthority)
533 st.encodeHeaderField(":path", "/")
534 return st.headerBuf.Bytes()
535 }
536
537 pseudoCount := map[string]int{}
538 keys := []string{":method", ":scheme", ":authority", ":path"}
539 vals := map[string][]string{
540 ":method": {"GET"},
541 ":scheme": {"https"},
542 ":authority": {defaultAuthority},
543 ":path": {"/"},
544 }
545 for len(headers) > 0 {
546 k, v := headers[0], headers[1]
547 headers = headers[2:]
548 if _, ok := vals[k]; !ok {
549 keys = append(keys, k)
550 }
551 if strings.HasPrefix(k, ":") {
552 pseudoCount[k]++
553 if pseudoCount[k] == 1 {
554 vals[k] = []string{v}
555 } else {
556
557 vals[k] = append(vals[k], v)
558 }
559 } else {
560 vals[k] = append(vals[k], v)
561 }
562 }
563 for _, k := range keys {
564 for _, v := range vals[k] {
565 st.encodeHeaderField(k, v)
566 }
567 }
568 return st.headerBuf.Bytes()
569 }
570
571
572 func (st *serverTester) bodylessReq1(headers ...string) {
573 st.writeHeaders(HeadersFrameParam{
574 StreamID: 1,
575 BlockFragment: st.encodeHeader(headers...),
576 EndStream: true,
577 EndHeaders: true,
578 })
579 }
580
581 func (st *serverTester) wantConnFlowControlConsumed(consumed int32) {
582 if got, want := st.sc.TestFlowControlConsumed(), consumed; got != want {
583 st.t.Errorf("connection flow control consumed: %v, want %v", got, want)
584 }
585 }
586
587 func TestServer(t *testing.T) { synctest.Test(t, testServer) }
588 func testServer(t *testing.T) {
589 gotReq := make(chan bool, 1)
590 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
591 w.Header().Set("Foo", "Bar")
592 gotReq <- true
593 })
594 defer st.Close()
595
596 st.greet()
597 st.writeHeaders(HeadersFrameParam{
598 StreamID: 1,
599 BlockFragment: st.encodeHeader(),
600 EndStream: true,
601 EndHeaders: true,
602 })
603
604 <-gotReq
605 }
606
607 func TestServer_Request_TLS(t *testing.T) {
608 for _, unencrypted := range []bool{false, true} {
609 for _, scheme := range []string{"https", "http", ""} {
610 name := scheme
611 if scheme == "" {
612 name = "CONNECT"
613 }
614 t.Run(fmt.Sprintf("unencrypted=%v/%s", unencrypted, name), func(t *testing.T) {
615 synctest.Test(t, func(t *testing.T) {
616 gotTLS := make(chan *tls.ConnectionState, 1)
617 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
618 gotTLS <- r.TLS
619 }, func(s *http.Server) {
620 s.Protocols = new(http.Protocols)
621 s.Protocols.SetHTTP2(!unencrypted)
622 s.Protocols.SetUnencryptedHTTP2(unencrypted)
623 })
624 st.greet()
625 headers := []string{":method", "CONNECT", ":authority", "example.com:443"}
626 if scheme != "" {
627 headers = []string{":method", "GET", ":authority", "example.com", ":scheme", scheme, ":path", "/"}
628 }
629 st.writeHeaders(HeadersFrameParam{
630 StreamID: 1,
631 BlockFragment: st.encodeHeaderRaw(headers...),
632 EndStream: true,
633 EndHeaders: true,
634 })
635 state := <-gotTLS
636 if unencrypted {
637 if state != nil {
638 t.Fatalf("Request.TLS = %v; want nil for an unencrypted connection", state)
639 }
640 } else {
641 if state == nil {
642 t.Fatal("Request.TLS = nil; want TLS connection state")
643 }
644 if !state.HandshakeComplete || state.NegotiatedProtocol != "h2" {
645 t.Errorf("Request.TLS = %+v; want completed HTTP/2 TLS handshake", state)
646 }
647 }
648 })
649 })
650 }
651 }
652 }
653
654 func TestServer_Request_Get(t *testing.T) { synctest.Test(t, testServer_Request_Get) }
655 func testServer_Request_Get(t *testing.T) {
656 testServerRequest(t, func(st *serverTester) {
657 st.writeHeaders(HeadersFrameParam{
658 StreamID: 1,
659 BlockFragment: st.encodeHeader("foo-bar", "some-value"),
660 EndStream: true,
661 EndHeaders: true,
662 })
663 }, func(r *http.Request) {
664 if r.Method != "GET" {
665 t.Errorf("Method = %q; want GET", r.Method)
666 }
667 if r.URL.Path != "/" {
668 t.Errorf("URL.Path = %q; want /", r.URL.Path)
669 }
670 if r.ContentLength != 0 {
671 t.Errorf("ContentLength = %v; want 0", r.ContentLength)
672 }
673 if r.Close {
674 t.Error("Close = true; want false")
675 }
676 if !strings.Contains(r.RemoteAddr, ":") {
677 t.Errorf("RemoteAddr = %q; want something with a colon", r.RemoteAddr)
678 }
679 if r.Proto != "HTTP/2.0" || r.ProtoMajor != 2 || r.ProtoMinor != 0 {
680 t.Errorf("Proto = %q Major=%v,Minor=%v; want HTTP/2.0", r.Proto, r.ProtoMajor, r.ProtoMinor)
681 }
682 wantHeader := http.Header{
683 "Foo-Bar": []string{"some-value"},
684 }
685 if !reflect.DeepEqual(r.Header, wantHeader) {
686 t.Errorf("Header = %#v; want %#v", r.Header, wantHeader)
687 }
688 if n, err := r.Body.Read([]byte(" ")); err != io.EOF || n != 0 {
689 t.Errorf("Read = %d, %v; want 0, EOF", n, err)
690 }
691 })
692 }
693
694 func TestServer_Request_Get_PathSlashes(t *testing.T) {
695 synctest.Test(t, testServer_Request_Get_PathSlashes)
696 }
697 func testServer_Request_Get_PathSlashes(t *testing.T) {
698 testServerRequest(t, func(st *serverTester) {
699 st.writeHeaders(HeadersFrameParam{
700 StreamID: 1,
701 BlockFragment: st.encodeHeader(":path", "/%2f/"),
702 EndStream: true,
703 EndHeaders: true,
704 })
705 }, func(r *http.Request) {
706 if r.RequestURI != "/%2f/" {
707 t.Errorf("RequestURI = %q; want /%%2f/", r.RequestURI)
708 }
709 if r.URL.Path != "///" {
710 t.Errorf("URL.Path = %q; want ///", r.URL.Path)
711 }
712 })
713 }
714
715
716
717
718
719 func TestServer_Request_Post_NoContentLength_EndStream(t *testing.T) {
720 synctest.Test(t, testServer_Request_Post_NoContentLength_EndStream)
721 }
722 func testServer_Request_Post_NoContentLength_EndStream(t *testing.T) {
723 testServerRequest(t, func(st *serverTester) {
724 st.writeHeaders(HeadersFrameParam{
725 StreamID: 1,
726 BlockFragment: st.encodeHeader(":method", "POST"),
727 EndStream: true,
728 EndHeaders: true,
729 })
730 }, func(r *http.Request) {
731 if r.Method != "POST" {
732 t.Errorf("Method = %q; want POST", r.Method)
733 }
734 if r.ContentLength != 0 {
735 t.Errorf("ContentLength = %v; want 0", r.ContentLength)
736 }
737 if n, err := r.Body.Read([]byte(" ")); err != io.EOF || n != 0 {
738 t.Errorf("Read = %d, %v; want 0, EOF", n, err)
739 }
740 })
741 }
742
743 func TestServer_Request_Post_Body_ImmediateEOF(t *testing.T) {
744 synctest.Test(t, testServer_Request_Post_Body_ImmediateEOF)
745 }
746 func testServer_Request_Post_Body_ImmediateEOF(t *testing.T) {
747 testBodyContents(t, -1, "", func(st *serverTester) {
748 st.writeHeaders(HeadersFrameParam{
749 StreamID: 1,
750 BlockFragment: st.encodeHeader(":method", "POST"),
751 EndStream: false,
752 EndHeaders: true,
753 })
754 st.writeData(1, true, nil)
755 })
756 }
757
758 func TestServer_Request_Post_Body_OneData(t *testing.T) {
759 synctest.Test(t, testServer_Request_Post_Body_OneData)
760 }
761 func testServer_Request_Post_Body_OneData(t *testing.T) {
762 const content = "Some content"
763 testBodyContents(t, -1, content, func(st *serverTester) {
764 st.writeHeaders(HeadersFrameParam{
765 StreamID: 1,
766 BlockFragment: st.encodeHeader(":method", "POST"),
767 EndStream: false,
768 EndHeaders: true,
769 })
770 st.writeData(1, true, []byte(content))
771 })
772 }
773
774 func TestServer_Request_Post_Body_TwoData(t *testing.T) {
775 synctest.Test(t, testServer_Request_Post_Body_TwoData)
776 }
777 func testServer_Request_Post_Body_TwoData(t *testing.T) {
778 const content = "Some content"
779 testBodyContents(t, -1, content, func(st *serverTester) {
780 st.writeHeaders(HeadersFrameParam{
781 StreamID: 1,
782 BlockFragment: st.encodeHeader(":method", "POST"),
783 EndStream: false,
784 EndHeaders: true,
785 })
786 st.writeData(1, false, []byte(content[:5]))
787 st.writeData(1, true, []byte(content[5:]))
788 })
789 }
790
791 func TestServer_Request_Post_Body_ContentLength_Correct(t *testing.T) {
792 synctest.Test(t, testServer_Request_Post_Body_ContentLength_Correct)
793 }
794 func testServer_Request_Post_Body_ContentLength_Correct(t *testing.T) {
795 const content = "Some content"
796 testBodyContents(t, int64(len(content)), content, func(st *serverTester) {
797 st.writeHeaders(HeadersFrameParam{
798 StreamID: 1,
799 BlockFragment: st.encodeHeader(
800 ":method", "POST",
801 "content-length", strconv.Itoa(len(content)),
802 ),
803 EndStream: false,
804 EndHeaders: true,
805 })
806 st.writeData(1, true, []byte(content))
807 })
808 }
809
810 func TestServer_Request_Post_Body_ContentLength_TooLarge(t *testing.T) {
811 synctest.Test(t, testServer_Request_Post_Body_ContentLength_TooLarge)
812 }
813 func testServer_Request_Post_Body_ContentLength_TooLarge(t *testing.T) {
814 testBodyContentsFail(t, 3, "request declared a Content-Length of 3 but only wrote 2 bytes",
815 func(st *serverTester) {
816 st.writeHeaders(HeadersFrameParam{
817 StreamID: 1,
818 BlockFragment: st.encodeHeader(
819 ":method", "POST",
820 "content-length", "3",
821 ),
822 EndStream: false,
823 EndHeaders: true,
824 })
825 st.writeData(1, true, []byte("12"))
826 })
827 }
828
829 func TestServer_Request_Post_Body_ContentLength_EndStream(t *testing.T) {
830 testRejectRequest(t, func(st *serverTester) {
831 st.writeHeaders(HeadersFrameParam{
832 StreamID: 1,
833 BlockFragment: st.encodeHeader(
834 ":method", "POST",
835 "content-length", "3",
836 ),
837 EndStream: true,
838 EndHeaders: true,
839 })
840 })
841 }
842
843 func TestServer_Request_Post_Body_ContentLength_TooSmall(t *testing.T) {
844 synctest.Test(t, testServer_Request_Post_Body_ContentLength_TooSmall)
845 }
846 func testServer_Request_Post_Body_ContentLength_TooSmall(t *testing.T) {
847 testBodyContentsFail(t, 4, "sender tried to send more than declared Content-Length of 4 bytes",
848 func(st *serverTester) {
849 st.writeHeaders(HeadersFrameParam{
850 StreamID: 1,
851 BlockFragment: st.encodeHeader(
852 ":method", "POST",
853 "content-length", "4",
854 ),
855 EndStream: false,
856 EndHeaders: true,
857 })
858 st.writeData(1, true, []byte("12345"))
859
860
861 st.wantRSTStream(1, ErrCodeProtocol)
862 st.wantConnFlowControlConsumed(0)
863 })
864 }
865
866 func testBodyContents(t *testing.T, wantContentLength int64, wantBody string, write func(st *serverTester)) {
867 testServerRequest(t, write, func(r *http.Request) {
868 if r.Method != "POST" {
869 t.Errorf("Method = %q; want POST", r.Method)
870 }
871 if r.ContentLength != wantContentLength {
872 t.Errorf("ContentLength = %v; want %d", r.ContentLength, wantContentLength)
873 }
874 all, err := io.ReadAll(r.Body)
875 if err != nil {
876 t.Fatal(err)
877 }
878 if string(all) != wantBody {
879 t.Errorf("Read = %q; want %q", all, wantBody)
880 }
881 if err := r.Body.Close(); err != nil {
882 t.Fatalf("Close: %v", err)
883 }
884 })
885 }
886
887 func testBodyContentsFail(t *testing.T, wantContentLength int64, wantReadError string, write func(st *serverTester)) {
888 testServerRequest(t, write, func(r *http.Request) {
889 if r.Method != "POST" {
890 t.Errorf("Method = %q; want POST", r.Method)
891 }
892 if r.ContentLength != wantContentLength {
893 t.Errorf("ContentLength = %v; want %d", r.ContentLength, wantContentLength)
894 }
895 all, err := io.ReadAll(r.Body)
896 if err == nil {
897 t.Fatalf("expected an error (%q) reading from the body. Successfully read %q instead.",
898 wantReadError, all)
899 }
900 if !strings.Contains(err.Error(), wantReadError) {
901 t.Fatalf("Body.Read = %v; want substring %q", err, wantReadError)
902 }
903 if err := r.Body.Close(); err != nil {
904 t.Fatalf("Close: %v", err)
905 }
906 })
907 }
908
909
910 func TestServer_Request_Get_Host(t *testing.T) { synctest.Test(t, testServer_Request_Get_Host) }
911 func testServer_Request_Get_Host(t *testing.T) {
912 const host = "example.com"
913 testServerRequest(t, func(st *serverTester) {
914 st.writeHeaders(HeadersFrameParam{
915 StreamID: 1,
916 BlockFragment: st.encodeHeader(":authority", "", "host", host),
917 EndStream: true,
918 EndHeaders: true,
919 })
920 }, func(r *http.Request) {
921 if r.Host != host {
922 t.Errorf("Host = %q; want %q", r.Host, host)
923 }
924 })
925 }
926
927
928 func TestServer_Request_Get_Authority(t *testing.T) {
929 synctest.Test(t, testServer_Request_Get_Authority)
930 }
931 func testServer_Request_Get_Authority(t *testing.T) {
932 const host = "example.com"
933 testServerRequest(t, func(st *serverTester) {
934 st.writeHeaders(HeadersFrameParam{
935 StreamID: 1,
936 BlockFragment: st.encodeHeader(":authority", host),
937 EndStream: true,
938 EndHeaders: true,
939 })
940 }, func(r *http.Request) {
941 if r.Host != host {
942 t.Errorf("Host = %q; want %q", r.Host, host)
943 }
944 })
945 }
946
947 func TestServer_Request_WithContinuation(t *testing.T) {
948 synctest.Test(t, testServer_Request_WithContinuation)
949 }
950 func testServer_Request_WithContinuation(t *testing.T) {
951 wantHeader := http.Header{
952 "Foo-One": []string{"value-one"},
953 "Foo-Two": []string{"value-two"},
954 "Foo-Three": []string{"value-three"},
955 }
956 testServerRequest(t, func(st *serverTester) {
957 fullHeaders := st.encodeHeader(
958 "foo-one", "value-one",
959 "foo-two", "value-two",
960 "foo-three", "value-three",
961 )
962 remain := fullHeaders
963 chunks := 0
964 for len(remain) > 0 {
965 const maxChunkSize = 5
966 chunk := remain
967 if len(chunk) > maxChunkSize {
968 chunk = chunk[:maxChunkSize]
969 }
970 remain = remain[len(chunk):]
971
972 if chunks == 0 {
973 st.writeHeaders(HeadersFrameParam{
974 StreamID: 1,
975 BlockFragment: chunk,
976 EndStream: true,
977 EndHeaders: false,
978 })
979 } else {
980 err := st.fr.WriteContinuation(1, len(remain) == 0, chunk)
981 if err != nil {
982 t.Fatal(err)
983 }
984 }
985 chunks++
986 }
987 if chunks < 2 {
988 t.Fatal("too few chunks")
989 }
990 }, func(r *http.Request) {
991 if !reflect.DeepEqual(r.Header, wantHeader) {
992 t.Errorf("Header = %#v; want %#v", r.Header, wantHeader)
993 }
994 })
995 }
996
997
998 func TestServer_Request_CookieConcat(t *testing.T) { synctest.Test(t, testServer_Request_CookieConcat) }
999 func testServer_Request_CookieConcat(t *testing.T) {
1000 const host = "example.com"
1001 testServerRequest(t, func(st *serverTester) {
1002 st.bodylessReq1(
1003 ":authority", host,
1004 "cookie", "a=b",
1005 "cookie", "c=d",
1006 "cookie", "e=f",
1007 )
1008 }, func(r *http.Request) {
1009 const want = "a=b; c=d; e=f"
1010 if got := r.Header.Get("Cookie"); got != want {
1011 t.Errorf("Cookie = %q; want %q", got, want)
1012 }
1013 })
1014 }
1015
1016 func TestServer_Request_Reject_CapitalHeader(t *testing.T) {
1017 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1("UPPER", "v") })
1018 }
1019
1020 func TestServer_Request_Reject_HeaderFieldNameColon(t *testing.T) {
1021 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1("has:colon", "v") })
1022 }
1023
1024 func TestServer_Request_Reject_HeaderFieldNameNULL(t *testing.T) {
1025 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1("has\x00null", "v") })
1026 }
1027
1028 func TestServer_Request_Reject_HeaderFieldNameEmpty(t *testing.T) {
1029 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1("", "v") })
1030 }
1031
1032 func TestServer_Request_Reject_HeaderFieldValueNewline(t *testing.T) {
1033 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1("foo", "has\nnewline") })
1034 }
1035
1036 func TestServer_Request_Reject_HeaderFieldValueCR(t *testing.T) {
1037 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1("foo", "has\rcarriage") })
1038 }
1039
1040 func TestServer_Request_Reject_HeaderFieldValueDEL(t *testing.T) {
1041 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1("foo", "has\x7fdel") })
1042 }
1043
1044 func TestServer_Request_Reject_Pseudo_Missing_method(t *testing.T) {
1045 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1(":method", "") })
1046 }
1047
1048 func TestServer_Request_Reject_Pseudo_ExactlyOne(t *testing.T) {
1049
1050
1051 testRejectRequest(t, func(st *serverTester) {
1052 st.addLogFilter("duplicate pseudo-header")
1053 st.bodylessReq1(":method", "GET", ":method", "POST")
1054 })
1055 }
1056
1057 func TestServer_Request_Reject_Pseudo_AfterRegular(t *testing.T) {
1058
1059
1060
1061
1062
1063
1064 testRejectRequest(t, func(st *serverTester) {
1065 st.addLogFilter("pseudo-header after regular header")
1066 var buf bytes.Buffer
1067 enc := hpack.NewEncoder(&buf)
1068 enc.WriteField(hpack.HeaderField{Name: ":method", Value: "GET"})
1069 enc.WriteField(hpack.HeaderField{Name: "regular", Value: "foobar"})
1070 enc.WriteField(hpack.HeaderField{Name: ":path", Value: "/"})
1071 enc.WriteField(hpack.HeaderField{Name: ":scheme", Value: "https"})
1072 st.writeHeaders(HeadersFrameParam{
1073 StreamID: 1,
1074 BlockFragment: buf.Bytes(),
1075 EndStream: true,
1076 EndHeaders: true,
1077 })
1078 })
1079 }
1080
1081 func TestServer_Request_Reject_Pseudo_Missing_path(t *testing.T) {
1082 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1(":path", "") })
1083 }
1084
1085 func TestServer_Request_Reject_Pseudo_Missing_scheme(t *testing.T) {
1086 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1(":scheme", "") })
1087 }
1088
1089 func TestServer_Request_Reject_Pseudo_scheme_invalid(t *testing.T) {
1090 testRejectRequest(t, func(st *serverTester) { st.bodylessReq1(":scheme", "bogus") })
1091 }
1092
1093 func TestServer_Request_Reject_Pseudo_Unknown(t *testing.T) {
1094 testRejectRequest(t, func(st *serverTester) {
1095 st.addLogFilter(`invalid pseudo-header ":unknown_thing"`)
1096 st.bodylessReq1(":unknown_thing", "")
1097 })
1098 }
1099
1100 func TestServer_Request_Reject_Authority_Userinfo(t *testing.T) {
1101
1102
1103
1104 testRejectRequest(t, func(st *serverTester) {
1105 var buf bytes.Buffer
1106 enc := hpack.NewEncoder(&buf)
1107 enc.WriteField(hpack.HeaderField{Name: ":authority", Value: "userinfo@example.tld"})
1108 enc.WriteField(hpack.HeaderField{Name: ":method", Value: "GET"})
1109 enc.WriteField(hpack.HeaderField{Name: ":path", Value: "/"})
1110 enc.WriteField(hpack.HeaderField{Name: ":scheme", Value: "https"})
1111 st.writeHeaders(HeadersFrameParam{
1112 StreamID: 1,
1113 BlockFragment: buf.Bytes(),
1114 EndStream: true,
1115 EndHeaders: true,
1116 })
1117 })
1118 }
1119
1120 func testRejectRequest(t *testing.T, send func(*serverTester)) {
1121 synctest.Test(t, func(t *testing.T) {
1122 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
1123 t.Error("server request made it to handler; should've been rejected")
1124 })
1125 defer st.Close()
1126
1127 st.greet()
1128 send(st)
1129 st.wantRSTStream(1, ErrCodeProtocol)
1130 })
1131 }
1132
1133 func newServerTesterForError(t *testing.T) *serverTester {
1134 t.Helper()
1135 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
1136 t.Error("server request made it to handler; should've been rejected")
1137 }, optQuiet)
1138 st.greet()
1139 return st
1140 }
1141
1142
1143
1144
1145 func TestRejectFrameOnIdle_WindowUpdate(t *testing.T) {
1146 synctest.Test(t, testRejectFrameOnIdle_WindowUpdate)
1147 }
1148 func testRejectFrameOnIdle_WindowUpdate(t *testing.T) {
1149 st := newServerTesterForError(t)
1150 st.fr.WriteWindowUpdate(123, 456)
1151 st.wantGoAway(123, ErrCodeProtocol)
1152 }
1153 func TestRejectFrameOnIdle_Data(t *testing.T) { synctest.Test(t, testRejectFrameOnIdle_Data) }
1154 func testRejectFrameOnIdle_Data(t *testing.T) {
1155 st := newServerTesterForError(t)
1156 st.fr.WriteData(123, true, nil)
1157 st.wantGoAway(123, ErrCodeProtocol)
1158 }
1159 func TestRejectFrameOnIdle_RSTStream(t *testing.T) { synctest.Test(t, testRejectFrameOnIdle_RSTStream) }
1160 func testRejectFrameOnIdle_RSTStream(t *testing.T) {
1161 st := newServerTesterForError(t)
1162 st.fr.WriteRSTStream(123, ErrCodeCancel)
1163 st.wantGoAway(123, ErrCodeProtocol)
1164 }
1165
1166 func TestServer_Request_Connect(t *testing.T) { synctest.Test(t, testServer_Request_Connect) }
1167 func testServer_Request_Connect(t *testing.T) {
1168 testServerRequest(t, func(st *serverTester) {
1169 st.writeHeaders(HeadersFrameParam{
1170 StreamID: 1,
1171 BlockFragment: st.encodeHeaderRaw(
1172 ":method", "CONNECT",
1173 ":authority", "example.com:123",
1174 ),
1175 EndStream: true,
1176 EndHeaders: true,
1177 })
1178 }, func(r *http.Request) {
1179 if g, w := r.Method, "CONNECT"; g != w {
1180 t.Errorf("Method = %q; want %q", g, w)
1181 }
1182 if g, w := r.RequestURI, "example.com:123"; g != w {
1183 t.Errorf("RequestURI = %q; want %q", g, w)
1184 }
1185 if g, w := r.URL.Host, "example.com:123"; g != w {
1186 t.Errorf("URL.Host = %q; want %q", g, w)
1187 }
1188 })
1189 }
1190
1191 func TestServer_Request_Connect_InvalidPath(t *testing.T) {
1192 synctest.Test(t, testServer_Request_Connect_InvalidPath)
1193 }
1194 func testServer_Request_Connect_InvalidPath(t *testing.T) {
1195 testServerRejectsStream(t, ErrCodeProtocol, func(st *serverTester) {
1196 st.writeHeaders(HeadersFrameParam{
1197 StreamID: 1,
1198 BlockFragment: st.encodeHeaderRaw(
1199 ":method", "CONNECT",
1200 ":authority", "example.com:123",
1201 ":path", "/bogus",
1202 ),
1203 EndStream: true,
1204 EndHeaders: true,
1205 })
1206 })
1207 }
1208
1209 func TestServer_Request_Connect_InvalidScheme(t *testing.T) {
1210 synctest.Test(t, testServer_Request_Connect_InvalidScheme)
1211 }
1212 func testServer_Request_Connect_InvalidScheme(t *testing.T) {
1213 testServerRejectsStream(t, ErrCodeProtocol, func(st *serverTester) {
1214 st.writeHeaders(HeadersFrameParam{
1215 StreamID: 1,
1216 BlockFragment: st.encodeHeaderRaw(
1217 ":method", "CONNECT",
1218 ":authority", "example.com:123",
1219 ":scheme", "https",
1220 ),
1221 EndStream: true,
1222 EndHeaders: true,
1223 })
1224 })
1225 }
1226
1227 func TestServer_Ping(t *testing.T) { synctest.Test(t, testServer_Ping) }
1228 func testServer_Ping(t *testing.T) {
1229 st := newServerTester(t, nil)
1230 defer st.Close()
1231 st.greet()
1232
1233
1234 ackPingData := [8]byte{1, 2, 4, 8, 16, 32, 64, 128}
1235 if err := st.fr.WritePing(true, ackPingData); err != nil {
1236 t.Fatal(err)
1237 }
1238
1239
1240 pingData := [8]byte{1, 2, 3, 4, 5, 6, 7, 8}
1241 if err := st.fr.WritePing(false, pingData); err != nil {
1242 t.Fatal(err)
1243 }
1244
1245 pf := readFrame[*PingFrame](t, st)
1246 if !pf.Flags.Has(FlagPingAck) {
1247 t.Error("response ping doesn't have ACK set")
1248 }
1249 if pf.Data != pingData {
1250 t.Errorf("response ping has data %q; want %q", pf.Data, pingData)
1251 }
1252 }
1253
1254 type filterListener struct {
1255 net.Listener
1256 accept func(conn net.Conn) (net.Conn, error)
1257 }
1258
1259 func (l *filterListener) Accept() (net.Conn, error) {
1260 c, err := l.Listener.Accept()
1261 if err != nil {
1262 return nil, err
1263 }
1264 return l.accept(c)
1265 }
1266
1267 func TestServer_MaxQueuedControlFrames(t *testing.T) {
1268 synctest.Test(t, testServer_MaxQueuedControlFrames)
1269 }
1270 func testServer_MaxQueuedControlFrames(t *testing.T) {
1271
1272 DisableGoroutineTracking(t)
1273
1274 st := newServerTester(t, nil)
1275 st.greet()
1276
1277 st.cc.(*tls.Conn).NetConn().(*synctestNetConn).SetReadBufferSize(0)
1278
1279
1280
1281 const extraPings = 2
1282 for range MaxQueuedControlFrames + extraPings {
1283 pingData := [8]byte{1, 2, 3, 4, 5, 6, 7, 8}
1284 st.fr.WritePing(false, pingData)
1285 }
1286 synctest.Wait()
1287
1288
1289
1290 st.cc.(*tls.Conn).NetConn().(*synctestNetConn).SetReadBufferSize(math.MaxInt)
1291
1292 st.advance(GoAwayTimeout)
1293
1294 for range 10 {
1295 if st.readFrame() == nil {
1296 break
1297 }
1298 }
1299 st.wantClosed()
1300 }
1301
1302 func TestServer_RejectsLargeFrames(t *testing.T) { synctest.Test(t, testServer_RejectsLargeFrames) }
1303 func testServer_RejectsLargeFrames(t *testing.T) {
1304 if runtime.GOOS == "windows" || runtime.GOOS == "plan9" || runtime.GOOS == "zos" {
1305 t.Skip("see golang.org/issue/13434, golang.org/issue/37321")
1306 }
1307 st := newServerTester(t, nil)
1308 defer st.Close()
1309 st.greet()
1310
1311
1312
1313
1314 st.fr.WriteRawFrame(0xff, 0, 0, make([]byte, DefaultMaxReadFrameSize+1))
1315
1316 st.wantGoAway(0, ErrCodeFrameSize)
1317 st.advance(GoAwayTimeout)
1318 st.wantClosed()
1319 }
1320
1321 func TestServer_Handler_Sends_WindowUpdate(t *testing.T) {
1322 synctest.Test(t, testServer_Handler_Sends_WindowUpdate)
1323 }
1324 func testServer_Handler_Sends_WindowUpdate(t *testing.T) {
1325
1326
1327
1328
1329 const windowSize = 65535 * 2
1330 st := newServerTester(t, nil, func(h2 *http.HTTP2Config) {
1331 h2.MaxReceiveBufferPerConnection = windowSize
1332 h2.MaxReceiveBufferPerStream = windowSize
1333 })
1334 defer st.Close()
1335
1336 st.greet()
1337 st.writeHeaders(HeadersFrameParam{
1338 StreamID: 1,
1339 BlockFragment: st.encodeHeader(":method", "POST"),
1340 EndStream: false,
1341 EndHeaders: true,
1342 })
1343 call := st.nextHandlerCall()
1344
1345
1346
1347
1348 data := make([]byte, windowSize)
1349 st.writeData(1, false, data[:1024])
1350 call.do(readBodyHandler(t, string(data[:1024])))
1351
1352
1353
1354 st.writeData(1, false, data[1024:])
1355 st.wantWindowUpdate(0, 1024)
1356 st.wantWindowUpdate(1, 1024)
1357
1358
1359 call.do(readBodyHandler(t, string(data[1024:])))
1360 st.wantWindowUpdate(0, windowSize-1024)
1361 st.wantWindowUpdate(1, windowSize-1024)
1362 }
1363
1364
1365
1366 func TestServer_Handler_Sends_WindowUpdate_Padding(t *testing.T) {
1367 synctest.Test(t, testServer_Handler_Sends_WindowUpdate_Padding)
1368 }
1369 func testServer_Handler_Sends_WindowUpdate_Padding(t *testing.T) {
1370 const windowSize = 65535 * 2
1371 st := newServerTester(t, nil, func(h2 *http.HTTP2Config) {
1372 h2.MaxReceiveBufferPerConnection = windowSize
1373 h2.MaxReceiveBufferPerStream = windowSize
1374 })
1375 defer st.Close()
1376
1377 st.greet()
1378 st.writeHeaders(HeadersFrameParam{
1379 StreamID: 1,
1380 BlockFragment: st.encodeHeader(":method", "POST"),
1381 EndStream: false,
1382 EndHeaders: true,
1383 })
1384 call := st.nextHandlerCall()
1385
1386
1387
1388
1389 data := make([]byte, windowSize/2)
1390 pad := make([]byte, 4)
1391 st.writeDataPadded(1, false, data, pad)
1392
1393
1394
1395
1396 call.do(readBodyHandler(t, string(data)))
1397 st.wantWindowUpdate(0, uint32(len(data)+1+len(pad)))
1398 st.wantWindowUpdate(1, uint32(len(data)+1+len(pad)))
1399 }
1400
1401 func TestServer_Send_GoAway_After_Bogus_WindowUpdate(t *testing.T) {
1402 synctest.Test(t, testServer_Send_GoAway_After_Bogus_WindowUpdate)
1403 }
1404 func testServer_Send_GoAway_After_Bogus_WindowUpdate(t *testing.T) {
1405 st := newServerTester(t, nil)
1406 defer st.Close()
1407 st.greet()
1408 if err := st.fr.WriteWindowUpdate(0, 1<<31-1); err != nil {
1409 t.Fatal(err)
1410 }
1411 st.wantGoAway(0, ErrCodeFlowControl)
1412 }
1413
1414 func TestServer_Send_RstStream_After_Bogus_WindowUpdate(t *testing.T) {
1415 synctest.Test(t, testServer_Send_RstStream_After_Bogus_WindowUpdate)
1416 }
1417 func testServer_Send_RstStream_After_Bogus_WindowUpdate(t *testing.T) {
1418 inHandler := make(chan bool)
1419 blockHandler := make(chan bool)
1420 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
1421 inHandler <- true
1422 <-blockHandler
1423 })
1424 defer st.Close()
1425 defer close(blockHandler)
1426 st.greet()
1427 st.writeHeaders(HeadersFrameParam{
1428 StreamID: 1,
1429 BlockFragment: st.encodeHeader(":method", "POST"),
1430 EndStream: false,
1431 EndHeaders: true,
1432 })
1433 <-inHandler
1434
1435 if err := st.fr.WriteWindowUpdate(1, 1<<31-1); err != nil {
1436 t.Fatal(err)
1437 }
1438 st.wantRSTStream(1, ErrCodeFlowControl)
1439 }
1440
1441
1442
1443
1444 func testServerPostUnblock(t *testing.T,
1445 handler func(http.ResponseWriter, *http.Request) error,
1446 fn func(*serverTester),
1447 checkErr func(error),
1448 otherHeaders ...string) {
1449 inHandler := make(chan bool)
1450 errc := make(chan error, 1)
1451 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
1452 inHandler <- true
1453 errc <- handler(w, r)
1454 })
1455 defer st.Close()
1456 st.greet()
1457 st.writeHeaders(HeadersFrameParam{
1458 StreamID: 1,
1459 BlockFragment: st.encodeHeader(append([]string{":method", "POST"}, otherHeaders...)...),
1460 EndStream: false,
1461 EndHeaders: true,
1462 })
1463 <-inHandler
1464 fn(st)
1465 err := <-errc
1466 if checkErr != nil {
1467 checkErr(err)
1468 }
1469 }
1470
1471 func TestServer_RSTStream_Unblocks_Read(t *testing.T) {
1472 synctest.Test(t, testServer_RSTStream_Unblocks_Read)
1473 }
1474 func testServer_RSTStream_Unblocks_Read(t *testing.T) {
1475 testServerPostUnblock(t,
1476 func(w http.ResponseWriter, r *http.Request) (err error) {
1477 _, err = r.Body.Read(make([]byte, 1))
1478 return
1479 },
1480 func(st *serverTester) {
1481 if err := st.fr.WriteRSTStream(1, ErrCodeCancel); err != nil {
1482 t.Fatal(err)
1483 }
1484 },
1485 func(err error) {
1486 want := StreamError{StreamID: 0x1, Code: 0x8}
1487 if !reflect.DeepEqual(err, want) {
1488 t.Errorf("Read error = %v; want %v", err, want)
1489 }
1490 },
1491 )
1492 }
1493
1494 func TestServer_RSTStream_Unblocks_Header_Write(t *testing.T) {
1495
1496
1497 n := 50
1498 if testing.Short() {
1499 n = 5
1500 }
1501 for i := 0; i < n; i++ {
1502 synctest.Test(t, testServer_RSTStream_Unblocks_Header_Write)
1503 }
1504 }
1505
1506 func testServer_RSTStream_Unblocks_Header_Write(t *testing.T) {
1507 inHandler := make(chan bool, 1)
1508 unblockHandler := make(chan bool, 1)
1509 headerWritten := make(chan bool, 1)
1510 wroteRST := make(chan bool, 1)
1511
1512 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
1513 inHandler <- true
1514 <-wroteRST
1515 w.Header().Set("foo", "bar")
1516 w.WriteHeader(200)
1517 w.(http.Flusher).Flush()
1518 headerWritten <- true
1519 <-unblockHandler
1520 })
1521 defer st.Close()
1522
1523 st.greet()
1524 st.writeHeaders(HeadersFrameParam{
1525 StreamID: 1,
1526 BlockFragment: st.encodeHeader(":method", "POST"),
1527 EndStream: false,
1528 EndHeaders: true,
1529 })
1530 <-inHandler
1531 if err := st.fr.WriteRSTStream(1, ErrCodeCancel); err != nil {
1532 t.Fatal(err)
1533 }
1534 wroteRST <- true
1535 synctest.Wait()
1536 <-headerWritten
1537 unblockHandler <- true
1538 }
1539
1540 func TestServer_DeadConn_Unblocks_Read(t *testing.T) {
1541 synctest.Test(t, testServer_DeadConn_Unblocks_Read)
1542 }
1543 func testServer_DeadConn_Unblocks_Read(t *testing.T) {
1544 testServerPostUnblock(t,
1545 func(w http.ResponseWriter, r *http.Request) (err error) {
1546 _, err = r.Body.Read(make([]byte, 1))
1547 return
1548 },
1549 func(st *serverTester) { st.cc.Close() },
1550 func(err error) {
1551 if err == nil {
1552 t.Error("unexpected nil error from Request.Body.Read")
1553 }
1554 },
1555 )
1556 }
1557
1558 var blockUntilClosed = func(w http.ResponseWriter, r *http.Request) error {
1559 <-w.(http.CloseNotifier).CloseNotify()
1560 return nil
1561 }
1562
1563 func TestServer_CloseNotify_After_RSTStream(t *testing.T) {
1564 synctest.Test(t, testServer_CloseNotify_After_RSTStream)
1565 }
1566 func testServer_CloseNotify_After_RSTStream(t *testing.T) {
1567 testServerPostUnblock(t, blockUntilClosed, func(st *serverTester) {
1568 if err := st.fr.WriteRSTStream(1, ErrCodeCancel); err != nil {
1569 t.Fatal(err)
1570 }
1571 }, nil)
1572 }
1573
1574 func TestServer_CloseNotify_After_ConnClose(t *testing.T) {
1575 synctest.Test(t, testServer_CloseNotify_After_ConnClose)
1576 }
1577 func testServer_CloseNotify_After_ConnClose(t *testing.T) {
1578 testServerPostUnblock(t, blockUntilClosed, func(st *serverTester) { st.cc.Close() }, nil)
1579 }
1580
1581
1582
1583
1584 func TestServer_CloseNotify_After_StreamError(t *testing.T) {
1585 synctest.Test(t, testServer_CloseNotify_After_StreamError)
1586 }
1587 func testServer_CloseNotify_After_StreamError(t *testing.T) {
1588 testServerPostUnblock(t, blockUntilClosed, func(st *serverTester) {
1589
1590 st.writeData(1, true, []byte("1234"))
1591 }, nil, "content-length", "3")
1592 }
1593
1594 func TestServer_StateTransitions(t *testing.T) { synctest.Test(t, testServer_StateTransitions) }
1595 func testServer_StateTransitions(t *testing.T) {
1596 var st *serverTester
1597 inHandler := make(chan bool)
1598 writeData := make(chan bool)
1599 leaveHandler := make(chan bool)
1600 st = newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
1601 inHandler <- true
1602 if !st.streamExists(1) {
1603 t.Errorf("stream 1 does not exist in handler")
1604 }
1605 if got, want := st.streamState(1), StateOpen; got != want {
1606 t.Errorf("in handler, state is %v; want %v", got, want)
1607 }
1608 writeData <- true
1609 if n, err := r.Body.Read(make([]byte, 1)); n != 0 || err != io.EOF {
1610 t.Errorf("body read = %d, %v; want 0, EOF", n, err)
1611 }
1612 if got, want := st.streamState(1), StateHalfClosedRemote; got != want {
1613 t.Errorf("in handler, state is %v; want %v", got, want)
1614 }
1615
1616 <-leaveHandler
1617 })
1618 st.greet()
1619 if st.streamExists(1) {
1620 t.Fatal("stream 1 should be empty")
1621 }
1622 if got := st.streamState(1); got != StateIdle {
1623 t.Fatalf("stream 1 should be idle; got %v", got)
1624 }
1625
1626 st.writeHeaders(HeadersFrameParam{
1627 StreamID: 1,
1628 BlockFragment: st.encodeHeader(":method", "POST"),
1629 EndStream: false,
1630 EndHeaders: true,
1631 })
1632 <-inHandler
1633 <-writeData
1634 st.writeData(1, true, nil)
1635
1636 leaveHandler <- true
1637 st.wantHeaders(wantHeader{
1638 streamID: 1,
1639 endStream: true,
1640 })
1641
1642 if got, want := st.streamState(1), StateClosed; got != want {
1643 t.Errorf("at end, state is %v; want %v", got, want)
1644 }
1645 if st.streamExists(1) {
1646 t.Fatal("at end, stream 1 should be gone")
1647 }
1648 }
1649
1650
1651 func TestServer_Rejects_HeadersNoEnd_Then_Headers(t *testing.T) {
1652 synctest.Test(t, testServer_Rejects_HeadersNoEnd_Then_Headers)
1653 }
1654 func testServer_Rejects_HeadersNoEnd_Then_Headers(t *testing.T) {
1655 st := newServerTesterForError(t)
1656 st.writeHeaders(HeadersFrameParam{
1657 StreamID: 1,
1658 BlockFragment: st.encodeHeader(),
1659 EndStream: true,
1660 EndHeaders: false,
1661 })
1662 st.writeHeaders(HeadersFrameParam{
1663 StreamID: 3,
1664 BlockFragment: st.encodeHeader(),
1665 EndStream: true,
1666 EndHeaders: true,
1667 })
1668 st.wantGoAway(0, ErrCodeProtocol)
1669 }
1670
1671
1672 func TestServer_Rejects_HeadersNoEnd_Then_Ping(t *testing.T) {
1673 synctest.Test(t, testServer_Rejects_HeadersNoEnd_Then_Ping)
1674 }
1675 func testServer_Rejects_HeadersNoEnd_Then_Ping(t *testing.T) {
1676 st := newServerTesterForError(t)
1677 st.writeHeaders(HeadersFrameParam{
1678 StreamID: 1,
1679 BlockFragment: st.encodeHeader(),
1680 EndStream: true,
1681 EndHeaders: false,
1682 })
1683 if err := st.fr.WritePing(false, [8]byte{}); err != nil {
1684 t.Fatal(err)
1685 }
1686 st.wantGoAway(0, ErrCodeProtocol)
1687 }
1688
1689
1690 func TestServer_Rejects_HeadersEnd_Then_Continuation(t *testing.T) {
1691 synctest.Test(t, testServer_Rejects_HeadersEnd_Then_Continuation)
1692 }
1693 func testServer_Rejects_HeadersEnd_Then_Continuation(t *testing.T) {
1694 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {}, optQuiet)
1695 st.greet()
1696 st.writeHeaders(HeadersFrameParam{
1697 StreamID: 1,
1698 BlockFragment: st.encodeHeader(),
1699 EndStream: true,
1700 EndHeaders: true,
1701 })
1702 st.wantHeaders(wantHeader{
1703 streamID: 1,
1704 endStream: true,
1705 })
1706 if err := st.fr.WriteContinuation(1, true, EncodeHeaderRaw(t, "foo", "bar")); err != nil {
1707 t.Fatal(err)
1708 }
1709 st.wantGoAway(1, ErrCodeProtocol)
1710 }
1711
1712
1713 func TestServer_Rejects_HeadersNoEnd_Then_ContinuationWrongStream(t *testing.T) {
1714 synctest.Test(t, testServer_Rejects_HeadersNoEnd_Then_ContinuationWrongStream)
1715 }
1716 func testServer_Rejects_HeadersNoEnd_Then_ContinuationWrongStream(t *testing.T) {
1717 st := newServerTesterForError(t)
1718 st.writeHeaders(HeadersFrameParam{
1719 StreamID: 1,
1720 BlockFragment: st.encodeHeader(),
1721 EndStream: true,
1722 EndHeaders: false,
1723 })
1724 if err := st.fr.WriteContinuation(3, true, EncodeHeaderRaw(t, "foo", "bar")); err != nil {
1725 t.Fatal(err)
1726 }
1727 st.wantGoAway(0, ErrCodeProtocol)
1728 }
1729
1730
1731 func TestServer_Rejects_Headers0(t *testing.T) { synctest.Test(t, testServer_Rejects_Headers0) }
1732 func testServer_Rejects_Headers0(t *testing.T) {
1733 st := newServerTesterForError(t)
1734 st.fr.AllowIllegalWrites = true
1735 st.writeHeaders(HeadersFrameParam{
1736 StreamID: 0,
1737 BlockFragment: st.encodeHeader(),
1738 EndStream: true,
1739 EndHeaders: true,
1740 })
1741 st.wantGoAway(0, ErrCodeProtocol)
1742 }
1743
1744
1745 func TestServer_Rejects_Continuation0(t *testing.T) {
1746 synctest.Test(t, testServer_Rejects_Continuation0)
1747 }
1748 func testServer_Rejects_Continuation0(t *testing.T) {
1749 st := newServerTesterForError(t)
1750 st.fr.AllowIllegalWrites = true
1751 if err := st.fr.WriteContinuation(0, true, st.encodeHeader()); err != nil {
1752 t.Fatal(err)
1753 }
1754 st.wantGoAway(0, ErrCodeProtocol)
1755 }
1756
1757
1758 func TestServer_Rejects_Priority0(t *testing.T) { synctest.Test(t, testServer_Rejects_Priority0) }
1759 func testServer_Rejects_Priority0(t *testing.T) {
1760 st := newServerTesterForError(t)
1761 st.fr.AllowIllegalWrites = true
1762 st.writePriority(0, PriorityParam{StreamDep: 1})
1763 st.wantGoAway(0, ErrCodeProtocol)
1764 }
1765
1766
1767
1768 func TestServer_Rejects_PriorityUpdate0(t *testing.T) {
1769 synctest.Test(t, testServer_Rejects_PriorityUpdate0)
1770 }
1771 func testServer_Rejects_PriorityUpdate0(t *testing.T) {
1772 st := newServerTesterForError(t)
1773 st.fr.AllowIllegalWrites = true
1774 st.writePriorityUpdate(0, "")
1775 st.wantGoAway(0, ErrCodeProtocol)
1776 }
1777
1778
1779 func TestServer_Rejects_PriorityUpdateUnparsable(t *testing.T) {
1780 synctest.Test(t, testServer_Rejects_PriorityUnparsable)
1781 }
1782 func testServer_Rejects_PriorityUnparsable(t *testing.T) {
1783 st := newServerTester(t, nil)
1784 defer st.Close()
1785 st.greet()
1786 st.writePriorityUpdate(1, "Invalid dictionary: ((((")
1787 st.wantRSTStream(1, ErrCodeProtocol)
1788 }
1789
1790
1791 func TestServer_Rejects_HeadersSelfDependence(t *testing.T) {
1792 synctest.Test(t, testServer_Rejects_HeadersSelfDependence)
1793 }
1794 func testServer_Rejects_HeadersSelfDependence(t *testing.T) {
1795 testServerRejectsStream(t, ErrCodeProtocol, func(st *serverTester) {
1796 st.fr.AllowIllegalWrites = true
1797 st.writeHeaders(HeadersFrameParam{
1798 StreamID: 1,
1799 BlockFragment: st.encodeHeader(),
1800 EndStream: true,
1801 EndHeaders: true,
1802 Priority: PriorityParam{StreamDep: 1},
1803 })
1804 })
1805 }
1806
1807
1808 func TestServer_Rejects_PrioritySelfDependence(t *testing.T) {
1809 synctest.Test(t, testServer_Rejects_PrioritySelfDependence)
1810 }
1811 func testServer_Rejects_PrioritySelfDependence(t *testing.T) {
1812 testServerRejectsStream(t, ErrCodeProtocol, func(st *serverTester) {
1813 st.fr.AllowIllegalWrites = true
1814 st.writePriority(1, PriorityParam{StreamDep: 1})
1815 })
1816 }
1817
1818 func TestServer_Rejects_PushPromise(t *testing.T) { synctest.Test(t, testServer_Rejects_PushPromise) }
1819 func testServer_Rejects_PushPromise(t *testing.T) {
1820 st := newServerTesterForError(t)
1821 pp := PushPromiseParam{
1822 StreamID: 1,
1823 PromiseID: 3,
1824 }
1825 if err := st.fr.WritePushPromise(pp); err != nil {
1826 t.Fatal(err)
1827 }
1828 st.wantGoAway(1, ErrCodeProtocol)
1829 }
1830
1831
1832
1833 func testServerRejectsStream(t *testing.T, code ErrCode, writeReq func(*serverTester)) {
1834 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {})
1835 defer st.Close()
1836 st.greet()
1837 writeReq(st)
1838 st.wantRSTStream(1, code)
1839 }
1840
1841
1842
1843
1844 func testServerRequest(t *testing.T, writeReq func(*serverTester), checkReq func(*http.Request)) {
1845 gotReq := make(chan bool, 1)
1846 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
1847 if r.Body == nil {
1848 t.Fatal("nil Body")
1849 }
1850 checkReq(r)
1851 gotReq <- true
1852 })
1853 defer st.Close()
1854
1855 st.greet()
1856 writeReq(st)
1857 <-gotReq
1858 }
1859
1860 func getSlash(st *serverTester) { st.bodylessReq1() }
1861
1862 func TestServer_Response_NoData(t *testing.T) { synctest.Test(t, testServer_Response_NoData) }
1863 func testServer_Response_NoData(t *testing.T) {
1864 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
1865
1866 return nil
1867 }, func(st *serverTester) {
1868 getSlash(st)
1869 st.wantHeaders(wantHeader{
1870 streamID: 1,
1871 endStream: true,
1872 })
1873 })
1874 }
1875
1876 func TestServer_Response_NoData_Header_FooBar(t *testing.T) {
1877 synctest.Test(t, testServer_Response_NoData_Header_FooBar)
1878 }
1879 func testServer_Response_NoData_Header_FooBar(t *testing.T) {
1880 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
1881 w.Header().Set("Foo-Bar", "some-value")
1882 return nil
1883 }, func(st *serverTester) {
1884 getSlash(st)
1885 st.wantHeaders(wantHeader{
1886 streamID: 1,
1887 endStream: true,
1888 header: http.Header{
1889 ":status": []string{"200"},
1890 "foo-bar": []string{"some-value"},
1891 "content-length": []string{"0"},
1892 },
1893 })
1894 })
1895 }
1896
1897
1898
1899 func TestServerIgnoresContentLengthSignWhenWritingChunks(t *testing.T) {
1900 synctest.Test(t, testServerIgnoresContentLengthSignWhenWritingChunks)
1901 }
1902 func testServerIgnoresContentLengthSignWhenWritingChunks(t *testing.T) {
1903 tests := []struct {
1904 name string
1905 cl string
1906 wantCL string
1907 }{
1908 {
1909 name: "proper content-length",
1910 cl: "3",
1911 wantCL: "3",
1912 },
1913 {
1914 name: "ignore cl with plus sign",
1915 cl: "+3",
1916 wantCL: "0",
1917 },
1918 {
1919 name: "ignore cl with minus sign",
1920 cl: "-3",
1921 wantCL: "0",
1922 },
1923 {
1924 name: "max int64, for safe uint64->int64 conversion",
1925 cl: "9223372036854775807",
1926 wantCL: "9223372036854775807",
1927 },
1928 {
1929 name: "overflows int64, so ignored",
1930 cl: "9223372036854775808",
1931 wantCL: "0",
1932 },
1933 }
1934
1935 for _, tt := range tests {
1936 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
1937 w.Header().Set("content-length", tt.cl)
1938 return nil
1939 }, func(st *serverTester) {
1940 getSlash(st)
1941 st.wantHeaders(wantHeader{
1942 streamID: 1,
1943 endStream: true,
1944 header: http.Header{
1945 ":status": []string{"200"},
1946 "content-length": []string{tt.wantCL},
1947 },
1948 })
1949 })
1950 }
1951 }
1952
1953
1954
1955 func TestServerRejectsContentLengthWithSignNewRequests(t *testing.T) {
1956 tests := []struct {
1957 name string
1958 cl string
1959 wantCL int64
1960 }{
1961 {
1962 name: "proper content-length",
1963 cl: "3",
1964 wantCL: 3,
1965 },
1966 {
1967 name: "ignore cl with plus sign",
1968 cl: "+3",
1969 wantCL: 0,
1970 },
1971 {
1972 name: "ignore cl with minus sign",
1973 cl: "-3",
1974 wantCL: 0,
1975 },
1976 {
1977 name: "max int64, for safe uint64->int64 conversion",
1978 cl: "9223372036854775807",
1979 wantCL: 9223372036854775807,
1980 },
1981 {
1982 name: "overflows int64, so ignored",
1983 cl: "9223372036854775808",
1984 wantCL: 0,
1985 },
1986 }
1987
1988 for _, tt := range tests {
1989 synctestSubtest(t, tt.name, func(t *testing.T) {
1990 writeReq := func(st *serverTester) {
1991 st.writeHeaders(HeadersFrameParam{
1992 StreamID: 1,
1993 BlockFragment: st.encodeHeader("content-length", tt.cl),
1994 EndStream: false,
1995 EndHeaders: true,
1996 })
1997 st.writeData(1, false, []byte(""))
1998 }
1999 checkReq := func(r *http.Request) {
2000 if r.ContentLength != tt.wantCL {
2001 t.Fatalf("Got: %d\nWant: %d", r.ContentLength, tt.wantCL)
2002 }
2003 }
2004 testServerRequest(t, writeReq, checkReq)
2005 })
2006 }
2007 }
2008
2009 func TestServerContentLengthDuplicates(t *testing.T) {
2010 tests := []struct {
2011 name string
2012 clValues []string
2013 wantOk bool
2014 }{
2015 {
2016 name: "single value",
2017 clValues: []string{"123"},
2018 wantOk: true,
2019 },
2020 {
2021 name: "identical duplicate values",
2022 clValues: []string{"123", "123", "123"},
2023 wantOk: true,
2024 },
2025 {
2026 name: "identical duplicate values with extra whitespace",
2027 clValues: []string{"123", " 123", "123"},
2028 wantOk: false,
2029 },
2030 {
2031 name: "different duplicate values",
2032 clValues: []string{"123", "321", "123"},
2033 wantOk: false,
2034 },
2035 }
2036 for _, tt := range tests {
2037 synctestSubtest(t, tt.name, func(t *testing.T) {
2038 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
2039 w.WriteHeader(200)
2040 })
2041 defer st.Close()
2042 st.greet()
2043
2044 headers := []string{":method", "GET"}
2045 for _, val := range tt.clValues {
2046 headers = append(headers, "content-length", val)
2047 }
2048 st.writeHeaders(HeadersFrameParam{
2049 StreamID: 1,
2050 BlockFragment: st.encodeHeader(headers...),
2051 EndStream: false,
2052 EndHeaders: true,
2053 })
2054 if tt.wantOk {
2055 st.wantHeaders(wantHeader{
2056 streamID: 1,
2057 endStream: true,
2058 header: http.Header{":status": []string{"200"}},
2059 })
2060 } else {
2061 st.wantRSTStream(1, ErrCodeProtocol)
2062 }
2063 })
2064 }
2065 }
2066
2067 func TestServer_Response_Data_Sniff_DoesntOverride(t *testing.T) {
2068 synctest.Test(t, testServer_Response_Data_Sniff_DoesntOverride)
2069 }
2070 func testServer_Response_Data_Sniff_DoesntOverride(t *testing.T) {
2071 const msg = "<html>this is HTML."
2072 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2073 w.Header().Set("Content-Type", "foo/bar")
2074 io.WriteString(w, msg)
2075 return nil
2076 }, func(st *serverTester) {
2077 getSlash(st)
2078 st.wantHeaders(wantHeader{
2079 streamID: 1,
2080 endStream: false,
2081 header: http.Header{
2082 ":status": []string{"200"},
2083 "content-type": []string{"foo/bar"},
2084 "content-length": []string{strconv.Itoa(len(msg))},
2085 },
2086 })
2087 st.wantData(wantData{
2088 streamID: 1,
2089 endStream: true,
2090 data: []byte(msg),
2091 })
2092 })
2093 }
2094
2095 func TestServer_Response_TransferEncoding_chunked(t *testing.T) {
2096 synctest.Test(t, testServer_Response_TransferEncoding_chunked)
2097 }
2098 func testServer_Response_TransferEncoding_chunked(t *testing.T) {
2099 const msg = "hi"
2100 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2101 w.Header().Set("Transfer-Encoding", "chunked")
2102 io.WriteString(w, msg)
2103 return nil
2104 }, func(st *serverTester) {
2105 getSlash(st)
2106 st.wantHeaders(wantHeader{
2107 streamID: 1,
2108 endStream: false,
2109 header: http.Header{
2110 ":status": []string{"200"},
2111 "content-type": []string{"text/plain; charset=utf-8"},
2112 "content-length": []string{strconv.Itoa(len(msg))},
2113 },
2114 })
2115 })
2116 }
2117
2118
2119 func TestServer_Response_Data_IgnoreHeaderAfterWrite_After(t *testing.T) {
2120 synctest.Test(t, testServer_Response_Data_IgnoreHeaderAfterWrite_After)
2121 }
2122 func testServer_Response_Data_IgnoreHeaderAfterWrite_After(t *testing.T) {
2123 const msg = "<html>this is HTML."
2124 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2125 io.WriteString(w, msg)
2126 w.Header().Set("foo", "should be ignored")
2127 return nil
2128 }, func(st *serverTester) {
2129 getSlash(st)
2130 st.wantHeaders(wantHeader{
2131 streamID: 1,
2132 endStream: false,
2133 header: http.Header{
2134 ":status": []string{"200"},
2135 "content-type": []string{"text/html; charset=utf-8"},
2136 "content-length": []string{strconv.Itoa(len(msg))},
2137 },
2138 })
2139 })
2140 }
2141
2142
2143 func TestServer_Response_Data_IgnoreHeaderAfterWrite_Overwrite(t *testing.T) {
2144 synctest.Test(t, testServer_Response_Data_IgnoreHeaderAfterWrite_Overwrite)
2145 }
2146 func testServer_Response_Data_IgnoreHeaderAfterWrite_Overwrite(t *testing.T) {
2147 const msg = "<html>this is HTML."
2148 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2149 w.Header().Set("foo", "proper value")
2150 io.WriteString(w, msg)
2151 w.Header().Set("foo", "should be ignored")
2152 return nil
2153 }, func(st *serverTester) {
2154 getSlash(st)
2155 st.wantHeaders(wantHeader{
2156 streamID: 1,
2157 endStream: false,
2158 header: http.Header{
2159 ":status": []string{"200"},
2160 "foo": []string{"proper value"},
2161 "content-type": []string{"text/html; charset=utf-8"},
2162 "content-length": []string{strconv.Itoa(len(msg))},
2163 },
2164 })
2165 })
2166 }
2167
2168 func TestServer_Response_Data_SniffLenType(t *testing.T) {
2169 synctest.Test(t, testServer_Response_Data_SniffLenType)
2170 }
2171 func testServer_Response_Data_SniffLenType(t *testing.T) {
2172 const msg = "<html>this is HTML."
2173 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2174 io.WriteString(w, msg)
2175 return nil
2176 }, func(st *serverTester) {
2177 getSlash(st)
2178 st.wantHeaders(wantHeader{
2179 streamID: 1,
2180 endStream: false,
2181 header: http.Header{
2182 ":status": []string{"200"},
2183 "content-type": []string{"text/html; charset=utf-8"},
2184 "content-length": []string{strconv.Itoa(len(msg))},
2185 },
2186 })
2187 st.wantData(wantData{
2188 streamID: 1,
2189 endStream: true,
2190 data: []byte(msg),
2191 })
2192 })
2193 }
2194
2195 func TestServer_Response_Header_Flush_MidWrite(t *testing.T) {
2196 synctest.Test(t, testServer_Response_Header_Flush_MidWrite)
2197 }
2198 func testServer_Response_Header_Flush_MidWrite(t *testing.T) {
2199 const msg = "<html>this is HTML"
2200 const msg2 = ", and this is the next chunk"
2201 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2202 io.WriteString(w, msg)
2203 w.(http.Flusher).Flush()
2204 io.WriteString(w, msg2)
2205 return nil
2206 }, func(st *serverTester) {
2207 getSlash(st)
2208 st.wantHeaders(wantHeader{
2209 streamID: 1,
2210 endStream: false,
2211 header: http.Header{
2212 ":status": []string{"200"},
2213 "content-type": []string{"text/html; charset=utf-8"},
2214
2215 },
2216 })
2217 st.wantData(wantData{
2218 streamID: 1,
2219 endStream: false,
2220 data: []byte(msg),
2221 })
2222 st.wantData(wantData{
2223 streamID: 1,
2224 endStream: true,
2225 data: []byte(msg2),
2226 })
2227 })
2228 }
2229
2230 func TestServer_Response_LargeWrite(t *testing.T) { synctest.Test(t, testServer_Response_LargeWrite) }
2231 func testServer_Response_LargeWrite(t *testing.T) {
2232 const size = 1 << 20
2233 const maxFrameSize = 16 << 10
2234 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2235 n, err := w.Write(bytes.Repeat([]byte("a"), size))
2236 if err != nil {
2237 return fmt.Errorf("Write error: %v", err)
2238 }
2239 if n != size {
2240 return fmt.Errorf("wrong size %d from Write", n)
2241 }
2242 return nil
2243 }, func(st *serverTester) {
2244 if err := st.fr.WriteSettings(
2245 Setting{SettingInitialWindowSize, 0},
2246 Setting{SettingMaxFrameSize, maxFrameSize},
2247 ); err != nil {
2248 t.Fatal(err)
2249 }
2250 st.wantSettingsAck()
2251
2252 getSlash(st)
2253
2254
2255 if err := st.fr.WriteWindowUpdate(1, size); err != nil {
2256 t.Fatal(err)
2257 }
2258
2259
2260 if err := st.fr.WriteWindowUpdate(0, size); err != nil {
2261 t.Fatal(err)
2262 }
2263 st.wantHeaders(wantHeader{
2264 streamID: 1,
2265 endStream: false,
2266 header: http.Header{
2267 ":status": []string{"200"},
2268 "content-type": []string{"text/plain; charset=utf-8"},
2269
2270 },
2271 })
2272 var bytes, frames int
2273 for {
2274 df := readFrame[*DataFrame](t, st)
2275 bytes += len(df.Data())
2276 frames++
2277 for _, b := range df.Data() {
2278 if b != 'a' {
2279 t.Fatal("non-'a' byte seen in DATA")
2280 }
2281 }
2282 if df.StreamEnded() {
2283 break
2284 }
2285 }
2286 if bytes != size {
2287 t.Errorf("Got %d bytes; want %d", bytes, size)
2288 }
2289 if want := int(size / maxFrameSize); frames < want || frames > want*2 {
2290 t.Errorf("Got %d frames; want %d", frames, size)
2291 }
2292 })
2293 }
2294
2295
2296 func TestServer_Response_LargeWrite_FlowControlled(t *testing.T) {
2297 synctest.Test(t, testServer_Response_LargeWrite_FlowControlled)
2298 }
2299 func testServer_Response_LargeWrite_FlowControlled(t *testing.T) {
2300
2301
2302 reads := []int{123, 1, 13, 127}
2303 size := 0
2304 for _, n := range reads {
2305 size += n
2306 }
2307
2308 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2309 w.(http.Flusher).Flush()
2310 n, err := w.Write(bytes.Repeat([]byte("a"), size))
2311 if err != nil {
2312 return fmt.Errorf("Write error: %v", err)
2313 }
2314 if n != size {
2315 return fmt.Errorf("wrong size %d from Write", n)
2316 }
2317 return nil
2318 }, func(st *serverTester) {
2319
2320
2321 if err := st.fr.WriteSettings(Setting{SettingInitialWindowSize, uint32(reads[0])}); err != nil {
2322 t.Fatal(err)
2323 }
2324 st.wantSettingsAck()
2325
2326 getSlash(st)
2327
2328 st.wantHeaders(wantHeader{
2329 streamID: 1,
2330 endStream: false,
2331 })
2332
2333 st.wantData(wantData{
2334 streamID: 1,
2335 endStream: false,
2336 size: reads[0],
2337 })
2338
2339 for i, quota := range reads[1:] {
2340 if err := st.fr.WriteWindowUpdate(1, uint32(quota)); err != nil {
2341 t.Fatal(err)
2342 }
2343 st.wantData(wantData{
2344 streamID: 1,
2345 endStream: i == len(reads[1:])-1,
2346 size: quota,
2347 })
2348 }
2349 })
2350 }
2351
2352
2353 func TestServer_Response_RST_Unblocks_LargeWrite(t *testing.T) {
2354 synctest.Test(t, testServer_Response_RST_Unblocks_LargeWrite)
2355 }
2356 func testServer_Response_RST_Unblocks_LargeWrite(t *testing.T) {
2357 const size = 1 << 20
2358 const maxFrameSize = 16 << 10
2359 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2360 w.(http.Flusher).Flush()
2361 _, err := w.Write(bytes.Repeat([]byte("a"), size))
2362 if err == nil {
2363 return errors.New("unexpected nil error from Write in handler")
2364 }
2365 return nil
2366 }, func(st *serverTester) {
2367 if err := st.fr.WriteSettings(
2368 Setting{SettingInitialWindowSize, 0},
2369 Setting{SettingMaxFrameSize, maxFrameSize},
2370 ); err != nil {
2371 t.Fatal(err)
2372 }
2373 st.wantSettingsAck()
2374
2375 getSlash(st)
2376
2377 st.wantHeaders(wantHeader{
2378 streamID: 1,
2379 endStream: false,
2380 })
2381
2382 if err := st.fr.WriteRSTStream(1, ErrCodeCancel); err != nil {
2383 t.Fatal(err)
2384 }
2385 })
2386 }
2387
2388 func TestServer_Response_Empty_Data_Not_FlowControlled(t *testing.T) {
2389 synctest.Test(t, testServer_Response_Empty_Data_Not_FlowControlled)
2390 }
2391 func testServer_Response_Empty_Data_Not_FlowControlled(t *testing.T) {
2392 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2393 w.(http.Flusher).Flush()
2394
2395 return nil
2396 }, func(st *serverTester) {
2397
2398 if err := st.fr.WriteSettings(Setting{SettingInitialWindowSize, 0}); err != nil {
2399 t.Fatal(err)
2400 }
2401 st.wantSettingsAck()
2402
2403 getSlash(st)
2404
2405 st.wantHeaders(wantHeader{
2406 streamID: 1,
2407 endStream: false,
2408 })
2409
2410 st.wantData(wantData{
2411 streamID: 1,
2412 endStream: true,
2413 size: 0,
2414 })
2415 })
2416 }
2417
2418 func TestServer_Response_Automatic100Continue(t *testing.T) {
2419 synctest.Test(t, testServer_Response_Automatic100Continue)
2420 }
2421 func testServer_Response_Automatic100Continue(t *testing.T) {
2422 const msg = "foo"
2423 const reply = "bar"
2424 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2425 if v := r.Header.Get("Expect"); v != "" {
2426 t.Errorf("Expect header = %q; want empty", v)
2427 }
2428 buf := make([]byte, len(msg))
2429
2430 if n, err := io.ReadFull(r.Body, buf); err != nil || n != len(msg) || string(buf) != msg {
2431 return fmt.Errorf("ReadFull = %q, %v; want %q, nil", buf[:n], err, msg)
2432 }
2433 _, err := io.WriteString(w, reply)
2434 return err
2435 }, func(st *serverTester) {
2436 st.writeHeaders(HeadersFrameParam{
2437 StreamID: 1,
2438 BlockFragment: st.encodeHeader(":method", "POST", "expect", "100-Continue"),
2439 EndStream: false,
2440 EndHeaders: true,
2441 })
2442 st.wantHeaders(wantHeader{
2443 streamID: 1,
2444 endStream: false,
2445 header: http.Header{
2446 ":status": []string{"100"},
2447 },
2448 })
2449
2450
2451
2452 st.writeData(1, true, []byte(msg))
2453
2454 st.wantHeaders(wantHeader{
2455 streamID: 1,
2456 endStream: false,
2457 header: http.Header{
2458 ":status": []string{"200"},
2459 "content-type": []string{"text/plain; charset=utf-8"},
2460 "content-length": []string{strconv.Itoa(len(reply))},
2461 },
2462 })
2463
2464 st.wantData(wantData{
2465 streamID: 1,
2466 endStream: true,
2467 data: []byte(reply),
2468 })
2469 })
2470 }
2471
2472 func TestServer_HandlerWriteErrorOnDisconnect(t *testing.T) {
2473 synctest.Test(t, testServer_HandlerWriteErrorOnDisconnect)
2474 }
2475 func testServer_HandlerWriteErrorOnDisconnect(t *testing.T) {
2476 errc := make(chan error, 1)
2477 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2478 p := []byte("some data.\n")
2479 for {
2480 _, err := w.Write(p)
2481 if err != nil {
2482 errc <- err
2483 return nil
2484 }
2485 }
2486 }, func(st *serverTester) {
2487 st.writeHeaders(HeadersFrameParam{
2488 StreamID: 1,
2489 BlockFragment: st.encodeHeader(),
2490 EndStream: false,
2491 EndHeaders: true,
2492 })
2493 st.wantHeaders(wantHeader{
2494 streamID: 1,
2495 endStream: false,
2496 })
2497
2498 st.cc.Close()
2499 _ = <-errc
2500 })
2501 }
2502
2503 func TestServer_Rejects_Too_Many_Streams(t *testing.T) {
2504 synctest.Test(t, testServer_Rejects_Too_Many_Streams)
2505 }
2506 func testServer_Rejects_Too_Many_Streams(t *testing.T) {
2507 st := newServerTester(t, nil)
2508 st.greet()
2509 nextStreamID := uint32(1)
2510 streamID := func() uint32 {
2511 defer func() { nextStreamID += 2 }()
2512 return nextStreamID
2513 }
2514 sendReq := func(id uint32) {
2515 st.writeHeaders(HeadersFrameParam{
2516 StreamID: id,
2517 BlockFragment: st.encodeHeader(
2518 ":path", fmt.Sprintf("/%v", id),
2519 ),
2520 EndStream: true,
2521 EndHeaders: true,
2522 })
2523 }
2524 var calls []*serverHandlerCall
2525 for range DefaultMaxStreams {
2526 sendReq(streamID())
2527 calls = append(calls, st.nextHandlerCall())
2528 }
2529
2530
2531
2532
2533 rejectID := streamID()
2534 headerBlock := st.encodeHeader(":path", fmt.Sprintf("/%v", rejectID))
2535 frag1, frag2 := headerBlock[:3], headerBlock[3:]
2536 st.writeHeaders(HeadersFrameParam{
2537 StreamID: rejectID,
2538 BlockFragment: frag1,
2539 EndStream: true,
2540 EndHeaders: false,
2541 })
2542 if err := st.fr.WriteContinuation(rejectID, true, frag2); err != nil {
2543 t.Fatal(err)
2544 }
2545 st.sync()
2546 st.wantRSTStream(rejectID, ErrCodeProtocol)
2547
2548
2549 calls[0].exit()
2550 st.sync()
2551 st.wantHeaders(wantHeader{
2552 streamID: 1,
2553 endStream: true,
2554 })
2555
2556
2557 goodID := streamID()
2558 sendReq(goodID)
2559 call := st.nextHandlerCall()
2560 if got, want := call.req.URL.Path, fmt.Sprintf("/%d", goodID); got != want {
2561 t.Errorf("Got request for %q, want %q", got, want)
2562 }
2563 }
2564
2565
2566 func TestServer_Response_ManyHeaders_With_Continuation(t *testing.T) {
2567 synctest.Test(t, testServer_Response_ManyHeaders_With_Continuation)
2568 }
2569 func testServer_Response_ManyHeaders_With_Continuation(t *testing.T) {
2570 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2571 h := w.Header()
2572 for i := range 5000 {
2573 h.Set(fmt.Sprintf("x-header-%d", i), fmt.Sprintf("x-value-%d", i))
2574 }
2575 return nil
2576 }, func(st *serverTester) {
2577 getSlash(st)
2578 hf := readFrame[*HeadersFrame](t, st)
2579 if hf.HeadersEnded() {
2580 t.Fatal("got unwanted END_HEADERS flag")
2581 }
2582 n := 0
2583 for {
2584 n++
2585 cf := readFrame[*ContinuationFrame](t, st)
2586 if cf.HeadersEnded() {
2587 break
2588 }
2589 }
2590 if n < 5 {
2591 t.Errorf("Only got %d CONTINUATION frames; expected 5+ (currently 6)", n)
2592 }
2593 })
2594 }
2595
2596
2597
2598
2599
2600
2601
2602
2603 func TestServer_NoCrash_HandlerClose_Then_ClientClose(t *testing.T) {
2604 synctest.Test(t, testServer_NoCrash_HandlerClose_Then_ClientClose)
2605 }
2606 func testServer_NoCrash_HandlerClose_Then_ClientClose(t *testing.T) {
2607 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2608
2609 return nil
2610 }, func(st *serverTester) {
2611 st.writeHeaders(HeadersFrameParam{
2612 StreamID: 1,
2613 BlockFragment: st.encodeHeader(),
2614 EndStream: false,
2615 EndHeaders: true,
2616 })
2617 st.wantHeaders(wantHeader{
2618 streamID: 1,
2619 endStream: true,
2620 })
2621
2622
2623
2624 st.wantRSTStream(1, ErrCodeNo)
2625
2626
2627
2628
2629
2630
2631 st.writeData(1, true, []byte("foo"))
2632
2633
2634
2635
2636
2637 st.wantRSTStream(1, ErrCodeStreamClosed)
2638
2639
2640
2641 st.wantConnFlowControlConsumed(0)
2642
2643
2644
2645 var (
2646 panMu sync.Mutex
2647 panicVal any
2648 )
2649
2650 SetTestHookOnPanic(t, func(sc *ServerConn, pv any) bool {
2651 panMu.Lock()
2652 panicVal = pv
2653 panMu.Unlock()
2654 return true
2655 })
2656
2657
2658 st.cc.Close()
2659 synctest.Wait()
2660
2661 panMu.Lock()
2662 got := panicVal
2663 panMu.Unlock()
2664 if got != nil {
2665 t.Errorf("Got panic: %v", got)
2666 }
2667 })
2668 }
2669
2670 func TestServer_Rejects_TLS10(t *testing.T) { testRejectTLS(t, tls.VersionTLS10) }
2671 func TestServer_Rejects_TLS11(t *testing.T) { testRejectTLS(t, tls.VersionTLS11) }
2672
2673 func testRejectTLS(t *testing.T, version uint16) {
2674 synctest.Test(t, func(t *testing.T) {
2675 st := newServerTester(t, nil, func(state *tls.ConnectionState) {
2676
2677
2678
2679 state.Version = version
2680 })
2681 defer st.Close()
2682 st.wantGoAway(0, ErrCodeInadequateSecurity)
2683 })
2684 }
2685
2686 func TestServer_Rejects_TLSBadCipher(t *testing.T) { synctest.Test(t, testServer_Rejects_TLSBadCipher) }
2687 func testServer_Rejects_TLSBadCipher(t *testing.T) {
2688 st := newServerTester(t, nil, func(state *tls.ConnectionState) {
2689 state.Version = tls.VersionTLS12
2690 state.CipherSuite = tls.TLS_RSA_WITH_RC4_128_SHA
2691 })
2692 defer st.Close()
2693 st.wantGoAway(0, ErrCodeInadequateSecurity)
2694 }
2695
2696 func TestServer_Advertises_Common_Cipher(t *testing.T) {
2697 synctest.Test(t, testServer_Advertises_Common_Cipher)
2698 }
2699 func testServer_Advertises_Common_Cipher(t *testing.T) {
2700 ts := newTestServer(t, func(w http.ResponseWriter, r *http.Request) {
2701 }, func(srv *http.Server) {
2702
2703
2704 srv.TLSConfig = nil
2705 })
2706
2707
2708 const requiredSuite = tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256
2709 tlsConfig := tlsConfigInsecure.Clone()
2710 tlsConfig.MaxVersion = tls.VersionTLS12
2711 tlsConfig.CipherSuites = []uint16{requiredSuite}
2712 tr := &http.Transport{
2713 TLSClientConfig: tlsConfig,
2714 Protocols: protocols("h2"),
2715 }
2716 defer tr.CloseIdleConnections()
2717
2718 req, err := http.NewRequest("GET", ts.URL, nil)
2719 if err != nil {
2720 t.Fatal(err)
2721 }
2722 res, err := tr.RoundTrip(req)
2723 if err != nil {
2724 t.Fatal(err)
2725 }
2726 res.Body.Close()
2727 }
2728
2729
2730
2731 func testServerResponse(t *testing.T,
2732 handler func(http.ResponseWriter, *http.Request) error,
2733 client func(*serverTester),
2734 ) {
2735 errc := make(chan error, 1)
2736 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
2737 if r.Body == nil {
2738 t.Fatal("nil Body")
2739 }
2740 err := handler(w, r)
2741 select {
2742 case errc <- err:
2743 default:
2744 t.Errorf("unexpected duplicate request")
2745 }
2746 })
2747 defer st.Close()
2748
2749 st.greet()
2750 client(st)
2751
2752 if err := <-errc; err != nil {
2753 t.Fatalf("Error in handler: %v", err)
2754 }
2755 }
2756
2757
2758
2759
2760 func readBodyHandler(t *testing.T, want string) func(w http.ResponseWriter, r *http.Request) {
2761 return func(w http.ResponseWriter, r *http.Request) {
2762 buf := make([]byte, len(want))
2763 _, err := io.ReadFull(r.Body, buf)
2764 if err != nil {
2765 t.Error(err)
2766 return
2767 }
2768 if string(buf) != want {
2769 t.Errorf("read %q; want %q", buf, want)
2770 }
2771 }
2772 }
2773
2774 func TestServer_MaxDecoderHeaderTableSize(t *testing.T) {
2775 synctest.Test(t, testServer_MaxDecoderHeaderTableSize)
2776 }
2777 func testServer_MaxDecoderHeaderTableSize(t *testing.T) {
2778 wantHeaderTableSize := uint32(InitialHeaderTableSize * 2)
2779 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {}, func(h2 *http.HTTP2Config) {
2780 h2.MaxDecoderHeaderTableSize = int(wantHeaderTableSize)
2781 })
2782 defer st.Close()
2783
2784 var advHeaderTableSize *uint32
2785 st.greetAndCheckSettings(func(s Setting) error {
2786 switch s.ID {
2787 case SettingHeaderTableSize:
2788 advHeaderTableSize = &s.Val
2789 }
2790 return nil
2791 })
2792
2793 if advHeaderTableSize == nil {
2794 t.Errorf("server didn't advertise a header table size")
2795 } else if got, want := *advHeaderTableSize, wantHeaderTableSize; got != want {
2796 t.Errorf("server advertised a header table size of %d, want %d", got, want)
2797 }
2798 }
2799
2800 func TestServer_MaxEncoderHeaderTableSize(t *testing.T) {
2801 synctest.Test(t, testServer_MaxEncoderHeaderTableSize)
2802 }
2803 func testServer_MaxEncoderHeaderTableSize(t *testing.T) {
2804 wantHeaderTableSize := uint32(InitialHeaderTableSize / 2)
2805 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {}, func(h2 *http.HTTP2Config) {
2806 h2.MaxEncoderHeaderTableSize = int(wantHeaderTableSize)
2807 })
2808 defer st.Close()
2809
2810 st.greet()
2811
2812 if got, want := st.sc.TestHPACKEncoder().MaxDynamicTableSize(), wantHeaderTableSize; got != want {
2813 t.Errorf("server encoder is using a header table size of %d, want %d", got, want)
2814 }
2815 }
2816
2817
2818 func TestServerDoS_MaxHeaderListSize(t *testing.T) { synctest.Test(t, testServerDoS_MaxHeaderListSize) }
2819 func testServerDoS_MaxHeaderListSize(t *testing.T) {
2820 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {})
2821 defer st.Close()
2822
2823
2824 frameSize := DefaultMaxReadFrameSize
2825 var advHeaderListSize *uint32
2826 st.greetAndCheckSettings(func(s Setting) error {
2827 switch s.ID {
2828 case SettingMaxFrameSize:
2829 if s.Val < MinMaxFrameSize {
2830 frameSize = MinMaxFrameSize
2831 } else if s.Val > MaxFrameSize {
2832 frameSize = MaxFrameSize
2833 } else {
2834 frameSize = int(s.Val)
2835 }
2836 case SettingMaxHeaderListSize:
2837 advHeaderListSize = &s.Val
2838 }
2839 return nil
2840 })
2841
2842 if advHeaderListSize == nil {
2843 t.Errorf("server didn't advertise a max header list size")
2844 } else if *advHeaderListSize == 0 {
2845 t.Errorf("server advertised a max header list size of 0")
2846 }
2847
2848 st.encodeHeaderField(":method", "GET")
2849 st.encodeHeaderField(":path", "/")
2850 st.encodeHeaderField(":scheme", "https")
2851 cookie := strings.Repeat("*", 4058)
2852 st.encodeHeaderField("cookie", cookie)
2853 st.writeHeaders(HeadersFrameParam{
2854 StreamID: 1,
2855 BlockFragment: st.headerBuf.Bytes(),
2856 EndStream: true,
2857 EndHeaders: false,
2858 })
2859
2860
2861
2862 st.headerBuf.Reset()
2863 st.encodeHeaderField("cookie", cookie)
2864
2865
2866 const size = 1 << 20
2867 b := bytes.Repeat(st.headerBuf.Bytes(), size/st.headerBuf.Len())
2868 for len(b) > 0 {
2869 chunk := b
2870 if len(chunk) > frameSize {
2871 chunk = chunk[:frameSize]
2872 }
2873 b = b[len(chunk):]
2874 st.fr.WriteContinuation(1, len(b) == 0, chunk)
2875 }
2876
2877 st.wantHeaders(wantHeader{
2878 streamID: 1,
2879 endStream: false,
2880 header: http.Header{
2881 ":status": []string{"431"},
2882 "content-type": []string{"text/html; charset=utf-8"},
2883 "content-length": []string{"63"},
2884 },
2885 })
2886 }
2887
2888 func TestServer_Response_Stream_With_Missing_Trailer(t *testing.T) {
2889 synctest.Test(t, testServer_Response_Stream_With_Missing_Trailer)
2890 }
2891 func testServer_Response_Stream_With_Missing_Trailer(t *testing.T) {
2892 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
2893 w.Header().Set("Trailer", "test-trailer")
2894 return nil
2895 }, func(st *serverTester) {
2896 getSlash(st)
2897 st.wantHeaders(wantHeader{
2898 streamID: 1,
2899 endStream: false,
2900 })
2901 st.wantData(wantData{
2902 streamID: 1,
2903 endStream: true,
2904 size: 0,
2905 })
2906 })
2907 }
2908
2909 func TestCompressionErrorOnWrite(t *testing.T) { synctest.Test(t, testCompressionErrorOnWrite) }
2910 func testCompressionErrorOnWrite(t *testing.T) {
2911 const maxStrLen = 8 << 10
2912 var serverConfig *http.Server
2913 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
2914
2915 }, func(s *http.Server) {
2916 serverConfig = s
2917 serverConfig.MaxHeaderBytes = maxStrLen
2918 })
2919 st.addLogFilter("connection error: COMPRESSION_ERROR")
2920 defer st.Close()
2921 st.greet()
2922
2923 maxAllowed := st.sc.TestFramerMaxHeaderStringLen()
2924
2925
2926
2927
2928
2929
2930 serverConfig.MaxHeaderBytes = 1 << 20
2931
2932
2933
2934
2935
2936 hbf := st.encodeHeader("foo", strings.Repeat("a", maxAllowed))
2937
2938 st.writeHeaders(HeadersFrameParam{
2939 StreamID: 1,
2940 BlockFragment: hbf,
2941 EndStream: true,
2942 EndHeaders: true,
2943 })
2944 st.wantHeaders(wantHeader{
2945 streamID: 1,
2946 endStream: false,
2947 header: http.Header{
2948 ":status": []string{"431"},
2949 "content-type": []string{"text/html; charset=utf-8"},
2950 "content-length": []string{"63"},
2951 },
2952 })
2953 df := readFrame[*DataFrame](t, st)
2954 if !strings.Contains(string(df.Data()), "HTTP Error 431") {
2955 t.Errorf("Unexpected data body: %q", df.Data())
2956 }
2957 if !df.StreamEnded() {
2958 t.Fatalf("expect data stream end")
2959 }
2960
2961
2962 hbf = st.encodeHeader("bar", strings.Repeat("b", maxAllowed+1))
2963 st.writeHeaders(HeadersFrameParam{
2964 StreamID: 3,
2965 BlockFragment: hbf,
2966 EndStream: true,
2967 EndHeaders: true,
2968 })
2969 st.wantGoAway(3, ErrCodeCompression)
2970 }
2971
2972 func TestCompressionErrorOnClose(t *testing.T) { synctest.Test(t, testCompressionErrorOnClose) }
2973 func testCompressionErrorOnClose(t *testing.T) {
2974 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
2975
2976 })
2977 st.addLogFilter("connection error: COMPRESSION_ERROR")
2978 defer st.Close()
2979 st.greet()
2980
2981 hbf := st.encodeHeader("foo", "bar")
2982 hbf = hbf[:len(hbf)-1]
2983 st.writeHeaders(HeadersFrameParam{
2984 StreamID: 1,
2985 BlockFragment: hbf,
2986 EndStream: true,
2987 EndHeaders: true,
2988 })
2989 st.wantGoAway(1, ErrCodeCompression)
2990 }
2991
2992
2993 func TestServerReadsTrailers(t *testing.T) { synctest.Test(t, testServerReadsTrailers) }
2994 func testServerReadsTrailers(t *testing.T) {
2995 const testBody = "some test body"
2996 writeReq := func(st *serverTester) {
2997 st.writeHeaders(HeadersFrameParam{
2998 StreamID: 1,
2999 BlockFragment: st.encodeHeader("trailer", "Foo, Bar", "trailer", "Baz"),
3000 EndStream: false,
3001 EndHeaders: true,
3002 })
3003 st.writeData(1, false, []byte(testBody))
3004 st.writeHeaders(HeadersFrameParam{
3005 StreamID: 1,
3006 BlockFragment: st.encodeHeaderRaw(
3007 "foo", "foov",
3008 "bar", "barv",
3009 "baz", "bazv",
3010 "surprise", "wasn't declared; shouldn't show up",
3011 ),
3012 EndStream: true,
3013 EndHeaders: true,
3014 })
3015 }
3016 checkReq := func(r *http.Request) {
3017 wantTrailer := http.Header{
3018 "Foo": nil,
3019 "Bar": nil,
3020 "Baz": nil,
3021 }
3022 if !reflect.DeepEqual(r.Trailer, wantTrailer) {
3023 t.Errorf("initial Trailer = %v; want %v", r.Trailer, wantTrailer)
3024 }
3025 slurp, err := io.ReadAll(r.Body)
3026 if string(slurp) != testBody {
3027 t.Errorf("read body %q; want %q", slurp, testBody)
3028 }
3029 if err != nil {
3030 t.Fatalf("Body slurp: %v", err)
3031 }
3032 wantTrailerAfter := http.Header{
3033 "Foo": {"foov"},
3034 "Bar": {"barv"},
3035 "Baz": {"bazv"},
3036 }
3037 if !reflect.DeepEqual(r.Trailer, wantTrailerAfter) {
3038 t.Errorf("final Trailer = %v; want %v", r.Trailer, wantTrailerAfter)
3039 }
3040 }
3041 testServerRequest(t, writeReq, checkReq)
3042 }
3043
3044
3045 func TestServerWritesTrailers_WithFlush(t *testing.T) {
3046 synctest.Test(t, func(t *testing.T) {
3047 testServerWritesTrailers(t, true)
3048 })
3049 }
3050 func TestServerWritesTrailers_WithoutFlush(t *testing.T) {
3051 synctest.Test(t, func(t *testing.T) {
3052 testServerWritesTrailers(t, false)
3053 })
3054 }
3055
3056 func testServerWritesTrailers(t *testing.T, withFlush bool) {
3057
3058 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
3059 w.Header().Set("Trailer", "Server-Trailer-A, Server-Trailer-B")
3060 w.Header().Add("Trailer", "Server-Trailer-C")
3061 w.Header().Add("Trailer", "Transfer-Encoding, Content-Length, Trailer")
3062
3063
3064 w.Header().Set("Foo", "Bar")
3065 w.Header().Set("Content-Length", "5")
3066
3067 io.WriteString(w, "Hello")
3068 if withFlush {
3069 w.(http.Flusher).Flush()
3070 }
3071 w.Header().Set("Server-Trailer-A", "valuea")
3072 w.Header().Set("Server-Trailer-C", "valuec")
3073
3074 w.Header().Set("Server-Surpise", "surprise! this isn't predeclared!")
3075
3076
3077
3078 w.Header().Set("Trailer:Post-Header-Trailer", "hi1")
3079 w.Header().Set("Trailer:post-header-trailer2", "hi2")
3080 w.Header().Set("Trailer:Range", "invalid")
3081 w.Header().Set("Trailer:Foo\x01Bogus", "invalid")
3082 w.Header().Set("Transfer-Encoding", "should not be included; Forbidden by RFC 7230 4.1.2")
3083 w.Header().Set("Content-Length", "should not be included; Forbidden by RFC 7230 4.1.2")
3084 w.Header().Set("Trailer", "should not be included; Forbidden by RFC 7230 4.1.2")
3085 return nil
3086 }, func(st *serverTester) {
3087
3088 st.h1server.ErrorLog = log.New(io.Discard, "", 0)
3089 getSlash(st)
3090 st.wantHeaders(wantHeader{
3091 streamID: 1,
3092 endStream: false,
3093 header: http.Header{
3094 ":status": []string{"200"},
3095 "foo": []string{"Bar"},
3096 "trailer": []string{
3097 "Server-Trailer-A, Server-Trailer-B",
3098 "Server-Trailer-C",
3099 "Transfer-Encoding, Content-Length, Trailer",
3100 },
3101 "content-type": []string{"text/plain; charset=utf-8"},
3102 "content-length": []string{"5"},
3103 },
3104 })
3105 st.wantData(wantData{
3106 streamID: 1,
3107 endStream: false,
3108 data: []byte("Hello"),
3109 })
3110 st.wantHeaders(wantHeader{
3111 streamID: 1,
3112 endStream: true,
3113 header: http.Header{
3114 "post-header-trailer": []string{"hi1"},
3115 "post-header-trailer2": []string{"hi2"},
3116 "server-trailer-a": []string{"valuea"},
3117 "server-trailer-c": []string{"valuec"},
3118 },
3119 })
3120 })
3121 }
3122
3123 func TestServerWritesUndeclaredTrailers(t *testing.T) {
3124 synctest.Test(t, testServerWritesUndeclaredTrailers)
3125 }
3126 func testServerWritesUndeclaredTrailers(t *testing.T) {
3127 const trailer = "Trailer-Header"
3128 const value = "hi1"
3129 ts := newTestServer(t, func(w http.ResponseWriter, r *http.Request) {
3130 w.Header().Set(http.TrailerPrefix+trailer, value)
3131 })
3132
3133 tr := &http.Transport{
3134 TLSClientConfig: tlsConfigInsecure,
3135 Protocols: protocols("h2"),
3136 }
3137 defer tr.CloseIdleConnections()
3138
3139 cl := &http.Client{Transport: tr}
3140 resp, err := cl.Get(ts.URL)
3141 if err != nil {
3142 t.Fatal(err)
3143 }
3144 io.Copy(io.Discard, resp.Body)
3145 resp.Body.Close()
3146
3147 if got, want := resp.Trailer.Get(trailer), value; got != want {
3148 t.Errorf("trailer %v = %q, want %q", trailer, got, want)
3149 }
3150 }
3151
3152
3153
3154 func TestServerDoesntWriteInvalidHeaders(t *testing.T) {
3155 synctest.Test(t, testServerDoesntWriteInvalidHeaders)
3156 }
3157 func testServerDoesntWriteInvalidHeaders(t *testing.T) {
3158 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
3159 w.Header().Add("OK1", "x")
3160 w.Header().Add("Bad:Colon", "x")
3161 w.Header().Add("Bad1\x00", "x")
3162 w.Header().Add("Bad2", "x\x00y")
3163 return nil
3164 }, func(st *serverTester) {
3165 getSlash(st)
3166 st.wantHeaders(wantHeader{
3167 streamID: 1,
3168 endStream: true,
3169 header: http.Header{
3170 ":status": []string{"200"},
3171 "ok1": []string{"x"},
3172 "content-length": []string{"0"},
3173 },
3174 })
3175 })
3176 }
3177
3178 func TestIssue53(t *testing.T) { synctest.Test(t, testIssue53) }
3179 func testIssue53(t *testing.T) {
3180 const data = "PRI * HTTP/2.0\r\n\r\nSM" +
3181 "\r\n\r\n\x00\x00\x00\x01\ainfinfin\ad"
3182 st := newServerTester(t, func(w http.ResponseWriter, req *http.Request) {
3183 w.Write([]byte("hello"))
3184 })
3185
3186 st.cc.Write([]byte(data))
3187 st.wantFrameType(FrameSettings)
3188 st.wantFrameType(FrameWindowUpdate)
3189 st.wantFrameType(FrameGoAway)
3190 time.Sleep(GoAwayTimeout)
3191 st.wantClosed()
3192 }
3193
3194 func TestServerServeNoBannedCiphers(t *testing.T) {
3195 tests := []struct {
3196 name string
3197 tlsConfig *tls.Config
3198 wantErr string
3199 }{
3200 {
3201 name: "empty CipherSuites",
3202 tlsConfig: &tls.Config{},
3203 },
3204 {
3205 name: "bad CipherSuites but MinVersion TLS 1.3",
3206 tlsConfig: &tls.Config{
3207 MinVersion: tls.VersionTLS13,
3208 CipherSuites: []uint16{tls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384},
3209 },
3210 },
3211 {
3212 name: "just the required cipher suite",
3213 tlsConfig: &tls.Config{
3214 CipherSuites: []uint16{tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256},
3215 },
3216 },
3217 {
3218 name: "just the alternative required cipher suite",
3219 tlsConfig: &tls.Config{
3220 CipherSuites: []uint16{tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256},
3221 },
3222 },
3223 {
3224 name: "missing required cipher suite",
3225 tlsConfig: &tls.Config{
3226 CipherSuites: []uint16{tls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384},
3227 },
3228 wantErr: "is missing an HTTP/2-required",
3229 },
3230 {
3231 name: "required after bad",
3232 tlsConfig: &tls.Config{
3233 CipherSuites: []uint16{tls.TLS_RSA_WITH_RC4_128_SHA, tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256},
3234 },
3235 },
3236 {
3237 name: "bad after required",
3238 tlsConfig: &tls.Config{
3239 CipherSuites: []uint16{tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256, tls.TLS_RSA_WITH_RC4_128_SHA},
3240 },
3241 },
3242 }
3243 for _, tt := range tests {
3244 tt.tlsConfig.Certificates = testServerTLSConfig.Certificates
3245
3246 srv := &http.Server{
3247 TLSConfig: tt.tlsConfig,
3248 Protocols: protocols("h2"),
3249 }
3250
3251 err := srv.ServeTLS(errListener{}, "", "")
3252 if (err != net.ErrClosed) != (tt.wantErr != "") {
3253 if tt.wantErr != "" {
3254 t.Errorf("%s: success, but want error", tt.name)
3255 } else {
3256 t.Errorf("%s: unexpected error: %v", tt.name, err)
3257 }
3258 }
3259 if err != nil && tt.wantErr != "" && !strings.Contains(err.Error(), tt.wantErr) {
3260 t.Errorf("%s: err = %v; want substring %q", tt.name, err, tt.wantErr)
3261 }
3262 if err == nil && !srv.TLSConfig.PreferServerCipherSuites {
3263 t.Errorf("%s: PreferServerCipherSuite is false; want true", tt.name)
3264 }
3265 }
3266 }
3267
3268 type errListener struct{}
3269
3270 func (li errListener) Accept() (net.Conn, error) { return nil, net.ErrClosed }
3271 func (li errListener) Close() error { return nil }
3272 func (li errListener) Addr() net.Addr { return nil }
3273
3274 func TestServerNoAutoContentLengthOnHead(t *testing.T) {
3275 synctest.Test(t, testServerNoAutoContentLengthOnHead)
3276 }
3277 func testServerNoAutoContentLengthOnHead(t *testing.T) {
3278 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3279
3280 })
3281 defer st.Close()
3282 st.greet()
3283 st.writeHeaders(HeadersFrameParam{
3284 StreamID: 1,
3285 BlockFragment: st.encodeHeader(":method", "HEAD"),
3286 EndStream: true,
3287 EndHeaders: true,
3288 })
3289 st.wantHeaders(wantHeader{
3290 streamID: 1,
3291 endStream: true,
3292 header: http.Header{
3293 ":status": []string{"200"},
3294 },
3295 })
3296 }
3297
3298
3299 func TestServerNoDuplicateContentType(t *testing.T) {
3300 synctest.Test(t, testServerNoDuplicateContentType)
3301 }
3302 func testServerNoDuplicateContentType(t *testing.T) {
3303 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3304 w.Header()["Content-Type"] = []string{""}
3305 fmt.Fprintf(w, "<html><head></head><body>hi</body></html>")
3306 })
3307 defer st.Close()
3308 st.greet()
3309 st.writeHeaders(HeadersFrameParam{
3310 StreamID: 1,
3311 BlockFragment: st.encodeHeader(),
3312 EndStream: true,
3313 EndHeaders: true,
3314 })
3315 st.wantHeaders(wantHeader{
3316 streamID: 1,
3317 endStream: false,
3318 header: http.Header{
3319 ":status": []string{"200"},
3320 "content-type": []string{""},
3321 "content-length": []string{"41"},
3322 },
3323 })
3324 }
3325
3326 func TestServerContentLengthCanBeDisabled(t *testing.T) {
3327 synctest.Test(t, testServerContentLengthCanBeDisabled)
3328 }
3329 func testServerContentLengthCanBeDisabled(t *testing.T) {
3330 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3331 w.Header()["Content-Length"] = nil
3332 fmt.Fprintf(w, "OK")
3333 })
3334 defer st.Close()
3335 st.greet()
3336 st.writeHeaders(HeadersFrameParam{
3337 StreamID: 1,
3338 BlockFragment: st.encodeHeader(),
3339 EndStream: true,
3340 EndHeaders: true,
3341 })
3342 st.wantHeaders(wantHeader{
3343 streamID: 1,
3344 endStream: false,
3345 header: http.Header{
3346 ":status": []string{"200"},
3347 "content-type": []string{"text/plain; charset=utf-8"},
3348 },
3349 })
3350 }
3351
3352
3353 func TestServer_Rejects_ConnHeaders(t *testing.T) { synctest.Test(t, testServer_Rejects_ConnHeaders) }
3354 func testServer_Rejects_ConnHeaders(t *testing.T) {
3355 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3356 t.Error("should not get to Handler")
3357 })
3358 defer st.Close()
3359 st.greet()
3360 st.bodylessReq1("connection", "foo")
3361 st.wantHeaders(wantHeader{
3362 streamID: 1,
3363 endStream: false,
3364 header: http.Header{
3365 ":status": []string{"400"},
3366 "content-type": []string{"text/plain; charset=utf-8"},
3367 "x-content-type-options": []string{"nosniff"},
3368 "content-length": []string{"51"},
3369 },
3370 })
3371 }
3372
3373 type hpackEncoder struct {
3374 enc *hpack.Encoder
3375 buf bytes.Buffer
3376 }
3377
3378 func (he *hpackEncoder) encodeHeaderRaw(t *testing.T, headers ...string) []byte {
3379 if len(headers)%2 == 1 {
3380 panic("odd number of kv args")
3381 }
3382 he.buf.Reset()
3383 if he.enc == nil {
3384 he.enc = hpack.NewEncoder(&he.buf)
3385 }
3386 for len(headers) > 0 {
3387 k, v := headers[0], headers[1]
3388 err := he.enc.WriteField(hpack.HeaderField{Name: k, Value: v})
3389 if err != nil {
3390 t.Fatalf("HPACK encoding error for %q/%q: %v", k, v, err)
3391 }
3392 headers = headers[2:]
3393 }
3394 return he.buf.Bytes()
3395 }
3396
3397
3398 func TestExpect100ContinueAfterHandlerWrites(t *testing.T) {
3399 synctest.Test(t, testExpect100ContinueAfterHandlerWrites)
3400 }
3401 func testExpect100ContinueAfterHandlerWrites(t *testing.T) {
3402 const msg = "Hello"
3403 const msg2 = "World"
3404
3405 doRead := make(chan bool, 1)
3406 defer close(doRead)
3407
3408 ts := newTestServer(t, func(w http.ResponseWriter, r *http.Request) {
3409 io.WriteString(w, msg)
3410 w.(http.Flusher).Flush()
3411
3412
3413 <-doRead
3414 r.Body.Read(make([]byte, 10))
3415
3416 io.WriteString(w, msg2)
3417 })
3418
3419 tr := &http.Transport{
3420 TLSClientConfig: tlsConfigInsecure,
3421 Protocols: protocols("h2"),
3422 }
3423 defer tr.CloseIdleConnections()
3424
3425 req, _ := http.NewRequest("POST", ts.URL, io.LimitReader(neverEnding('A'), 2<<20))
3426 req.Header.Set("Expect", "100-continue")
3427
3428 res, err := tr.RoundTrip(req)
3429 if err != nil {
3430 t.Fatal(err)
3431 }
3432 defer res.Body.Close()
3433
3434 buf := make([]byte, len(msg))
3435 if _, err := io.ReadFull(res.Body, buf); err != nil {
3436 t.Fatal(err)
3437 }
3438 if string(buf) != msg {
3439 t.Fatalf("msg = %q; want %q", buf, msg)
3440 }
3441
3442 doRead <- true
3443
3444 if _, err := io.ReadFull(res.Body, buf); err != nil {
3445 t.Fatal(err)
3446 }
3447 if string(buf) != msg2 {
3448 t.Fatalf("second msg = %q; want %q", buf, msg2)
3449 }
3450 }
3451
3452 type funcReader func([]byte) (n int, err error)
3453
3454 func (f funcReader) Read(p []byte) (n int, err error) { return f(p) }
3455
3456
3457
3458 func TestUnreadFlowControlReturned_Server(t *testing.T) {
3459 for _, tt := range []struct {
3460 name string
3461 reqFn func(r *http.Request)
3462 }{
3463 {
3464 "body-open",
3465 func(r *http.Request) {},
3466 },
3467 {
3468 "body-closed",
3469 func(r *http.Request) {
3470 r.Body.Close()
3471 },
3472 },
3473 {
3474 "read-1-byte-and-close",
3475 func(r *http.Request) {
3476 b := make([]byte, 1)
3477 r.Body.Read(b)
3478 r.Body.Close()
3479 },
3480 },
3481 } {
3482 synctestSubtest(t, tt.name, func(t *testing.T) {
3483 unblock := make(chan bool, 1)
3484 defer close(unblock)
3485
3486 ts := newTestServer(t, func(w http.ResponseWriter, r *http.Request) {
3487
3488
3489
3490 tt.reqFn(r)
3491 <-unblock
3492 })
3493
3494 tr := &http.Transport{
3495 TLSClientConfig: tlsConfigInsecure,
3496 Protocols: protocols("h2"),
3497 }
3498 defer tr.CloseIdleConnections()
3499
3500
3501 iters := 100
3502 if testing.Short() {
3503 iters = 20
3504 }
3505 for i := 0; i < iters; i++ {
3506 body := io.MultiReader(
3507 io.LimitReader(neverEnding('A'), 16<<10),
3508 funcReader(func([]byte) (n int, err error) {
3509 unblock <- true
3510 return 0, io.EOF
3511 }),
3512 )
3513 req, _ := http.NewRequest("POST", ts.URL, body)
3514 res, err := tr.RoundTrip(req)
3515 if err != nil {
3516 t.Fatal(tt.name, err)
3517 }
3518 res.Body.Close()
3519 }
3520 })
3521 }
3522 }
3523
3524 func TestServerReturnsStreamAndConnFlowControlOnBodyClose(t *testing.T) {
3525 synctest.Test(t, testServerReturnsStreamAndConnFlowControlOnBodyClose)
3526 }
3527 func testServerReturnsStreamAndConnFlowControlOnBodyClose(t *testing.T) {
3528 unblockHandler := make(chan struct{})
3529 defer close(unblockHandler)
3530
3531 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3532 r.Body.Close()
3533 w.WriteHeader(200)
3534 w.(http.Flusher).Flush()
3535 <-unblockHandler
3536 })
3537 defer st.Close()
3538
3539 st.greet()
3540 st.writeHeaders(HeadersFrameParam{
3541 StreamID: 1,
3542 BlockFragment: st.encodeHeader(),
3543 EndHeaders: true,
3544 })
3545 st.wantHeaders(wantHeader{
3546 streamID: 1,
3547 endStream: false,
3548 })
3549 const size = InflowMinRefresh
3550 st.writeData(1, false, make([]byte, size))
3551 st.wantWindowUpdate(0, size)
3552 unblockHandler <- struct{}{}
3553 st.wantData(wantData{
3554 streamID: 1,
3555 endStream: true,
3556 })
3557 }
3558
3559 func TestServerIdleTimeout(t *testing.T) { synctest.Test(t, testServerIdleTimeout) }
3560 func testServerIdleTimeout(t *testing.T) {
3561 if testing.Short() {
3562 t.Skip("skipping in short mode")
3563 }
3564
3565 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3566 }, func(s *http.Server) {
3567 s.IdleTimeout = 500 * time.Millisecond
3568 })
3569 defer st.Close()
3570
3571 st.greet()
3572 st.advance(500 * time.Millisecond)
3573 st.wantGoAway(0, ErrCodeNo)
3574 }
3575
3576 func TestServerIdleTimeout_AfterRequest(t *testing.T) {
3577 synctest.Test(t, testServerIdleTimeout_AfterRequest)
3578 }
3579 func testServerIdleTimeout_AfterRequest(t *testing.T) {
3580 if testing.Short() {
3581 t.Skip("skipping in short mode")
3582 }
3583 const (
3584 requestTimeout = 2 * time.Second
3585 idleTimeout = 1 * time.Second
3586 )
3587
3588 var st *serverTester
3589 st = newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3590 time.Sleep(requestTimeout)
3591 }, func(s *http.Server) {
3592 s.IdleTimeout = idleTimeout
3593 })
3594 defer st.Close()
3595
3596 st.greet()
3597
3598
3599
3600 st.bodylessReq1()
3601 st.advance(requestTimeout)
3602 st.wantHeaders(wantHeader{
3603 streamID: 1,
3604 endStream: true,
3605 })
3606
3607
3608
3609 st.advance(idleTimeout)
3610 st.wantGoAway(1, ErrCodeNo)
3611 }
3612
3613
3614
3615
3616 func TestRequestBodyReadCloseRace(t *testing.T) { synctest.Test(t, testRequestBodyReadCloseRace) }
3617 func testRequestBodyReadCloseRace(t *testing.T) {
3618 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3619 go r.Body.Close()
3620 io.Copy(io.Discard, r.Body)
3621 })
3622 st.greet()
3623
3624 data := make([]byte, 1024)
3625 for i := range 100 {
3626 streamID := uint32(1 + (i * 2))
3627 st.writeHeaders(HeadersFrameParam{
3628 StreamID: streamID,
3629 BlockFragment: st.encodeHeader(),
3630 EndHeaders: true,
3631 })
3632 st.writeData(1, false, data)
3633
3634 for {
3635
3636
3637 fr := st.readFrame()
3638 if fr == nil {
3639 t.Fatalf("got no RSTStreamFrame, want one")
3640 }
3641 rst, ok := fr.(*RSTStreamFrame)
3642 if !ok {
3643 continue
3644 }
3645
3646 if rst.ErrCode != ErrCodeNo && rst.ErrCode != ErrCodeStreamClosed {
3647 t.Fatalf("got RSTStreamFrame with error code %v, want ErrCodeNo or ErrCodeStreamClosed", rst.ErrCode)
3648 }
3649 break
3650 }
3651 }
3652 }
3653
3654 func TestIssue20704Race(t *testing.T) { synctest.Test(t, testIssue20704Race) }
3655 func testIssue20704Race(t *testing.T) {
3656 if testing.Short() && os.Getenv("GO_BUILDER_NAME") == "" {
3657 t.Skip("skipping in short mode")
3658 }
3659 const (
3660 itemSize = 1 << 10
3661 itemCount = 100
3662 )
3663
3664 ts := newTestServer(t, func(w http.ResponseWriter, r *http.Request) {
3665 for range itemCount {
3666 _, err := w.Write(make([]byte, itemSize))
3667 if err != nil {
3668 return
3669 }
3670 }
3671 })
3672
3673 tr := &http.Transport{
3674 TLSClientConfig: tlsConfigInsecure,
3675 Protocols: protocols("h2"),
3676 }
3677 defer tr.CloseIdleConnections()
3678 cl := &http.Client{Transport: tr}
3679
3680 for range 1000 {
3681 resp, err := cl.Get(ts.URL)
3682 if err != nil {
3683 t.Fatal(err)
3684 }
3685
3686
3687 resp.Body.Close()
3688 }
3689 }
3690
3691 func TestServer_Rejects_TooSmall(t *testing.T) { synctest.Test(t, testServer_Rejects_TooSmall) }
3692 func testServer_Rejects_TooSmall(t *testing.T) {
3693 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
3694 io.ReadAll(r.Body)
3695 return nil
3696 }, func(st *serverTester) {
3697 st.writeHeaders(HeadersFrameParam{
3698 StreamID: 1,
3699 BlockFragment: st.encodeHeader(
3700 ":method", "POST",
3701 "content-length", "4",
3702 ),
3703 EndStream: false,
3704 EndHeaders: true,
3705 })
3706 st.writeData(1, true, []byte("12345"))
3707 st.wantRSTStream(1, ErrCodeProtocol)
3708 st.wantConnFlowControlConsumed(0)
3709 })
3710 }
3711
3712
3713
3714 func TestServerHandlerConnectionClose(t *testing.T) {
3715 synctest.Test(t, testServerHandlerConnectionClose)
3716 }
3717 func testServerHandlerConnectionClose(t *testing.T) {
3718 unblockHandler := make(chan bool, 1)
3719 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
3720 w.Header().Set("Connection", "close")
3721 w.Header().Set("Foo", "bar")
3722 w.(http.Flusher).Flush()
3723 <-unblockHandler
3724 return nil
3725 }, func(st *serverTester) {
3726 defer close(unblockHandler)
3727 st.writeHeaders(HeadersFrameParam{
3728 StreamID: 1,
3729 BlockFragment: st.encodeHeader(),
3730 EndStream: true,
3731 EndHeaders: true,
3732 })
3733 var sawGoAway bool
3734 var sawRes bool
3735 var sawWindowUpdate bool
3736 for {
3737 f := st.readFrame()
3738 if f == nil {
3739 break
3740 }
3741 switch f := f.(type) {
3742 case *GoAwayFrame:
3743 sawGoAway = true
3744 if f.LastStreamID != 1 || f.ErrCode != ErrCodeNo {
3745 t.Errorf("unexpected GOAWAY frame: %v", SummarizeFrame(f))
3746 }
3747
3748
3749 st.writeHeaders(HeadersFrameParam{
3750 StreamID: 3,
3751 BlockFragment: st.encodeHeader(),
3752 EndStream: false,
3753 EndHeaders: true,
3754 })
3755 st.fr.WriteRSTStream(3, ErrCodeCancel)
3756
3757
3758
3759 st.writeHeaders(HeadersFrameParam{
3760 StreamID: 5,
3761 BlockFragment: st.encodeHeader(),
3762 EndStream: false,
3763 EndHeaders: true,
3764 })
3765
3766 st.writeData(5, true, make([]byte, 1<<19))
3767 case *HeadersFrame:
3768 goth := st.decodeHeader(f.HeaderBlockFragment())
3769 wanth := [][2]string{
3770 {":status", "200"},
3771 {"foo", "bar"},
3772 }
3773 if !reflect.DeepEqual(goth, wanth) {
3774 t.Errorf("got headers %v; want %v", goth, wanth)
3775 }
3776 sawRes = true
3777 case *DataFrame:
3778 if f.StreamID != 1 || !f.StreamEnded() || len(f.Data()) != 0 {
3779 t.Errorf("unexpected DATA frame: %v", SummarizeFrame(f))
3780 }
3781 case *WindowUpdateFrame:
3782 if !sawGoAway {
3783 t.Errorf("unexpected WINDOW_UPDATE frame: %v", SummarizeFrame(f))
3784 return
3785 }
3786 if f.StreamID != 0 {
3787 st.t.Fatalf("WindowUpdate StreamID = %d; want 5", f.FrameHeader.StreamID)
3788 return
3789 }
3790 sawWindowUpdate = true
3791 unblockHandler <- true
3792 st.sync()
3793 st.advance(GoAwayTimeout)
3794 default:
3795 t.Logf("unexpected frame: %v", SummarizeFrame(f))
3796 }
3797 }
3798 if !sawGoAway {
3799 t.Errorf("didn't see GOAWAY")
3800 }
3801 if !sawRes {
3802 t.Errorf("didn't see response")
3803 }
3804 if !sawWindowUpdate {
3805 t.Errorf("didn't see WINDOW_UPDATE")
3806 }
3807 })
3808 }
3809
3810 func TestServer_Headers_HalfCloseRemote(t *testing.T) {
3811 synctest.Test(t, testServer_Headers_HalfCloseRemote)
3812 }
3813 func testServer_Headers_HalfCloseRemote(t *testing.T) {
3814 var st *serverTester
3815 writeData := make(chan bool)
3816 writeHeaders := make(chan bool)
3817 leaveHandler := make(chan bool)
3818 st = newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3819 if !st.streamExists(1) {
3820 t.Errorf("stream 1 does not exist in handler")
3821 }
3822 if got, want := st.streamState(1), StateOpen; got != want {
3823 t.Errorf("in handler, state is %v; want %v", got, want)
3824 }
3825 writeData <- true
3826 if n, err := r.Body.Read(make([]byte, 1)); n != 0 || err != io.EOF {
3827 t.Errorf("body read = %d, %v; want 0, EOF", n, err)
3828 }
3829 if got, want := st.streamState(1), StateHalfClosedRemote; got != want {
3830 t.Errorf("in handler, state is %v; want %v", got, want)
3831 }
3832 writeHeaders <- true
3833
3834 <-leaveHandler
3835 })
3836 st.greet()
3837
3838 st.writeHeaders(HeadersFrameParam{
3839 StreamID: 1,
3840 BlockFragment: st.encodeHeader(),
3841 EndStream: false,
3842 EndHeaders: true,
3843 })
3844 <-writeData
3845 st.writeData(1, true, nil)
3846
3847 <-writeHeaders
3848
3849 st.writeHeaders(HeadersFrameParam{
3850 StreamID: 1,
3851 BlockFragment: st.encodeHeader(),
3852 EndStream: false,
3853 EndHeaders: true,
3854 })
3855
3856 defer close(leaveHandler)
3857
3858 st.wantRSTStream(1, ErrCodeStreamClosed)
3859 }
3860
3861 func TestServerGracefulShutdown(t *testing.T) { synctest.Test(t, testServerGracefulShutdown) }
3862 func testServerGracefulShutdown(t *testing.T) {
3863 handlerDone := make(chan struct{})
3864 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
3865 <-handlerDone
3866 w.Header().Set("x-foo", "bar")
3867 })
3868 defer st.Close()
3869
3870 st.greet()
3871 st.bodylessReq1()
3872
3873 st.sync()
3874
3875 shutdownc := make(chan struct{})
3876 go func() {
3877 defer close(shutdownc)
3878 st.h1server.Shutdown(context.Background())
3879 }()
3880
3881 st.wantGoAway(1, ErrCodeNo)
3882
3883 close(handlerDone)
3884 st.sync()
3885
3886 st.wantHeaders(wantHeader{
3887 streamID: 1,
3888 endStream: true,
3889 header: http.Header{
3890 ":status": []string{"200"},
3891 "x-foo": []string{"bar"},
3892 "content-length": []string{"0"},
3893 },
3894 })
3895
3896 n, err := st.cc.Read([]byte{0})
3897 if n != 0 || err == nil {
3898 t.Errorf("Read = %v, %v; want 0, non-nil", n, err)
3899 }
3900
3901
3902 <-shutdownc
3903 }
3904
3905
3906 func TestContentEncodingNoSniffing(t *testing.T) {
3907 type resp struct {
3908 name string
3909 body []byte
3910
3911
3912
3913 contentEncoding any
3914 wantContentType string
3915 }
3916
3917 resps := []*resp{
3918 {
3919 name: "gzip content-encoding, gzipped",
3920 contentEncoding: "application/gzip",
3921 wantContentType: "",
3922 body: func() []byte {
3923 buf := new(bytes.Buffer)
3924 gzw := gzip.NewWriter(buf)
3925 gzw.Write([]byte("doctype html><p>Hello</p>"))
3926 gzw.Close()
3927 return buf.Bytes()
3928 }(),
3929 },
3930 {
3931 name: "zlib content-encoding, zlibbed",
3932 contentEncoding: "application/zlib",
3933 wantContentType: "",
3934 body: func() []byte {
3935 buf := new(bytes.Buffer)
3936 zw := zlib.NewWriter(buf)
3937 zw.Write([]byte("doctype html><p>Hello</p>"))
3938 zw.Close()
3939 return buf.Bytes()
3940 }(),
3941 },
3942 {
3943 name: "no content-encoding",
3944 wantContentType: "application/x-gzip",
3945 body: func() []byte {
3946 buf := new(bytes.Buffer)
3947 gzw := gzip.NewWriter(buf)
3948 gzw.Write([]byte("doctype html><p>Hello</p>"))
3949 gzw.Close()
3950 return buf.Bytes()
3951 }(),
3952 },
3953 {
3954 name: "phony content-encoding",
3955 contentEncoding: "foo/bar",
3956 body: []byte("doctype html><p>Hello</p>"),
3957 },
3958 {
3959 name: "empty but set content-encoding",
3960 contentEncoding: "",
3961 wantContentType: "audio/mpeg",
3962 body: []byte("ID3"),
3963 },
3964 }
3965
3966 for _, tt := range resps {
3967 synctestSubtest(t, tt.name, func(t *testing.T) {
3968 ts := newTestServer(t, func(w http.ResponseWriter, r *http.Request) {
3969 if tt.contentEncoding != nil {
3970 w.Header().Set("Content-Encoding", tt.contentEncoding.(string))
3971 }
3972 w.Write(tt.body)
3973 })
3974
3975 tr := &http.Transport{
3976 TLSClientConfig: tlsConfigInsecure,
3977 Protocols: protocols("h2"),
3978 }
3979 defer tr.CloseIdleConnections()
3980
3981 req, _ := http.NewRequest("GET", ts.URL, nil)
3982 res, err := tr.RoundTrip(req)
3983 if err != nil {
3984 t.Fatalf("GET %s: %v", ts.URL, err)
3985 }
3986 defer res.Body.Close()
3987
3988 g := res.Header.Get("Content-Encoding")
3989 t.Logf("%s: Content-Encoding: %s", ts.URL, g)
3990
3991 if w := tt.contentEncoding; g != w {
3992 if w != nil {
3993 t.Errorf("Content-Encoding mismatch\n\tgot: %q\n\twant: %q", g, w)
3994 } else if g != "" {
3995 t.Errorf("Unexpected Content-Encoding %q", g)
3996 }
3997 }
3998
3999 g = res.Header.Get("Content-Type")
4000 if w := tt.wantContentType; g != w {
4001 t.Errorf("Content-Type mismatch\n\tgot: %q\n\twant: %q", g, w)
4002 }
4003 t.Logf("%s: Content-Type: %s", ts.URL, g)
4004 })
4005 }
4006 }
4007
4008 func TestServerWindowUpdateOnBodyClose(t *testing.T) {
4009 synctest.Test(t, testServerWindowUpdateOnBodyClose)
4010 }
4011 func testServerWindowUpdateOnBodyClose(t *testing.T) {
4012 const windowSize = 65535 * 2
4013 content := make([]byte, windowSize)
4014 errc := make(chan error)
4015 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4016 buf := make([]byte, 4)
4017 n, err := io.ReadFull(r.Body, buf)
4018 if err != nil {
4019 errc <- err
4020 return
4021 }
4022 if n != len(buf) {
4023 errc <- fmt.Errorf("too few bytes read: %d", n)
4024 return
4025 }
4026 r.Body.Close()
4027 errc <- nil
4028 }, func(h2 *http.HTTP2Config) {
4029 h2.MaxReceiveBufferPerConnection = windowSize
4030 h2.MaxReceiveBufferPerStream = windowSize
4031 })
4032 defer st.Close()
4033
4034 st.greet()
4035 st.writeHeaders(HeadersFrameParam{
4036 StreamID: 1,
4037 BlockFragment: st.encodeHeader(
4038 ":method", "POST",
4039 "content-length", strconv.Itoa(len(content)),
4040 ),
4041 EndStream: false,
4042 EndHeaders: true,
4043 })
4044 st.writeData(1, false, content[:windowSize/2])
4045 if err := <-errc; err != nil {
4046 t.Fatal(err)
4047 }
4048
4049
4050 increments := windowSize / 2
4051 for {
4052 f := st.readFrame()
4053 if f == nil {
4054 break
4055 }
4056 if wu, ok := f.(*WindowUpdateFrame); ok && wu.StreamID == 0 {
4057 increments -= int(wu.Increment)
4058 if increments == 0 {
4059 break
4060 }
4061 }
4062 }
4063
4064
4065 st.writeData(1, false, content[windowSize/2:])
4066 st.wantWindowUpdate(0, windowSize/2)
4067 }
4068
4069 func TestNoErrorLoggedOnPostAfterGOAWAY(t *testing.T) {
4070 synctest.Test(t, testNoErrorLoggedOnPostAfterGOAWAY)
4071 }
4072 func testNoErrorLoggedOnPostAfterGOAWAY(t *testing.T) {
4073 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {})
4074 defer st.Close()
4075
4076 st.greet()
4077
4078 content := "some content"
4079 st.writeHeaders(HeadersFrameParam{
4080 StreamID: 1,
4081 BlockFragment: st.encodeHeader(
4082 ":method", "POST",
4083 "content-length", strconv.Itoa(len(content)),
4084 ),
4085 EndStream: false,
4086 EndHeaders: true,
4087 })
4088 st.wantHeaders(wantHeader{
4089 streamID: 1,
4090 endStream: true,
4091 })
4092
4093 st.sc.StartGracefulShutdown()
4094 st.wantRSTStream(1, ErrCodeNo)
4095 st.wantGoAway(1, ErrCodeNo)
4096
4097 st.writeData(1, true, []byte(content))
4098 st.Close()
4099
4100 if bytes.Contains(st.serverLogBuf.Bytes(), []byte("PROTOCOL_ERROR")) {
4101 t.Error("got protocol error")
4102 }
4103 }
4104
4105 func TestServerSendsProcessing(t *testing.T) { synctest.Test(t, testServerSendsProcessing) }
4106 func testServerSendsProcessing(t *testing.T) {
4107 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
4108 w.WriteHeader(http.StatusProcessing)
4109 w.Write([]byte("stuff"))
4110
4111 return nil
4112 }, func(st *serverTester) {
4113 getSlash(st)
4114 st.wantHeaders(wantHeader{
4115 streamID: 1,
4116 endStream: false,
4117 header: http.Header{
4118 ":status": []string{"102"},
4119 },
4120 })
4121 st.wantHeaders(wantHeader{
4122 streamID: 1,
4123 endStream: false,
4124 header: http.Header{
4125 ":status": []string{"200"},
4126 "content-type": []string{"text/plain; charset=utf-8"},
4127 "content-length": []string{"5"},
4128 },
4129 })
4130 })
4131 }
4132
4133 func TestServerSendsEarlyHints(t *testing.T) { synctest.Test(t, testServerSendsEarlyHints) }
4134 func testServerSendsEarlyHints(t *testing.T) {
4135 testServerResponse(t, func(w http.ResponseWriter, r *http.Request) error {
4136 h := w.Header()
4137 h.Add("Content-Length", "123")
4138 h.Add("Link", "</style.css>; rel=preload; as=style")
4139 h.Add("Link", "</script.js>; rel=preload; as=script")
4140 w.WriteHeader(http.StatusEarlyHints)
4141
4142 h.Add("Link", "</foo.js>; rel=preload; as=script")
4143 w.WriteHeader(http.StatusEarlyHints)
4144
4145 w.Write([]byte("stuff"))
4146
4147 return nil
4148 }, func(st *serverTester) {
4149 getSlash(st)
4150 st.wantHeaders(wantHeader{
4151 streamID: 1,
4152 endStream: false,
4153 header: http.Header{
4154 ":status": []string{"103"},
4155 "link": []string{
4156 "</style.css>; rel=preload; as=style",
4157 "</script.js>; rel=preload; as=script",
4158 },
4159 },
4160 })
4161 st.wantHeaders(wantHeader{
4162 streamID: 1,
4163 endStream: false,
4164 header: http.Header{
4165 ":status": []string{"103"},
4166 "link": []string{
4167 "</style.css>; rel=preload; as=style",
4168 "</script.js>; rel=preload; as=script",
4169 "</foo.js>; rel=preload; as=script",
4170 },
4171 },
4172 })
4173 st.wantHeaders(wantHeader{
4174 streamID: 1,
4175 endStream: false,
4176 header: http.Header{
4177 ":status": []string{"200"},
4178 "link": []string{
4179 "</style.css>; rel=preload; as=style",
4180 "</script.js>; rel=preload; as=script",
4181 "</foo.js>; rel=preload; as=script",
4182 },
4183 "content-type": []string{"text/plain; charset=utf-8"},
4184 "content-length": []string{"123"},
4185 },
4186 })
4187 })
4188 }
4189
4190 func TestProtocolErrorAfterGoAway(t *testing.T) { synctest.Test(t, testProtocolErrorAfterGoAway) }
4191 func testProtocolErrorAfterGoAway(t *testing.T) {
4192 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4193 io.Copy(io.Discard, r.Body)
4194 })
4195 defer st.Close()
4196
4197 st.greet()
4198 content := "some content"
4199 st.writeHeaders(HeadersFrameParam{
4200 StreamID: 1,
4201 BlockFragment: st.encodeHeader(
4202 ":method", "POST",
4203 "content-length", strconv.Itoa(len(content)),
4204 ),
4205 EndStream: false,
4206 EndHeaders: true,
4207 })
4208 st.writeData(1, false, []byte(content[:5]))
4209
4210
4211
4212 if err := st.fr.WriteGoAway(1, ErrCodeNo, nil); err != nil {
4213 t.Fatal(err)
4214 }
4215 if err := st.fr.WriteWindowUpdate(0, 1<<31-1); err != nil {
4216 t.Fatal(err)
4217 }
4218
4219 st.advance(GoAwayTimeout)
4220 st.wantGoAway(1, ErrCodeNo)
4221 st.wantClosed()
4222 }
4223
4224 func TestServerInitialFlowControlWindow(t *testing.T) {
4225 for _, want := range []int32{
4226 65535,
4227 1 << 19,
4228 1 << 21,
4229
4230
4231
4232
4233
4234 65535 * 2,
4235 } {
4236 synctestSubtest(t, fmt.Sprint(want), func(t *testing.T) {
4237
4238 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4239 }, func(h2 *http.HTTP2Config) {
4240 h2.MaxReceiveBufferPerConnection = int(want)
4241 })
4242 st.writePreface()
4243 st.writeSettings()
4244 _ = readFrame[*SettingsFrame](t, st)
4245 st.writeSettingsAck()
4246 st.writeHeaders(HeadersFrameParam{
4247 StreamID: 1,
4248 BlockFragment: st.encodeHeader(),
4249 EndStream: true,
4250 EndHeaders: true,
4251 })
4252 window := 65535
4253 Frames:
4254 for {
4255 f := st.readFrame()
4256 switch f := f.(type) {
4257 case *WindowUpdateFrame:
4258 if f.FrameHeader.StreamID != 0 {
4259 t.Errorf("WindowUpdate StreamID = %d; want 0", f.FrameHeader.StreamID)
4260 return
4261 }
4262 window += int(f.Increment)
4263 case *HeadersFrame:
4264 break Frames
4265 case nil:
4266 break Frames
4267 default:
4268 }
4269 }
4270 if window != int(want) {
4271 t.Errorf("got initial flow control window = %v, want %v", window, want)
4272 }
4273 })
4274 }
4275 }
4276
4277
4278
4279
4280
4281
4282 func TestServerWriteDoesNotRetainBufferAfterReturn(t *testing.T) {
4283 synctest.Test(t, testServerWriteDoesNotRetainBufferAfterReturn)
4284 }
4285 func testServerWriteDoesNotRetainBufferAfterReturn(t *testing.T) {
4286 donec := make(chan struct{})
4287 ts := newTestServer(t, func(w http.ResponseWriter, r *http.Request) {
4288 defer close(donec)
4289 buf := make([]byte, 1<<20)
4290 var i byte
4291 for {
4292 i++
4293 _, err := w.Write(buf)
4294 for j := range buf {
4295 buf[j] = byte(i)
4296 }
4297 if err != nil {
4298 return
4299 }
4300 }
4301 })
4302
4303 tr := &http.Transport{
4304 TLSClientConfig: tlsConfigInsecure,
4305 Protocols: protocols("h2"),
4306 }
4307 defer tr.CloseIdleConnections()
4308
4309 req, _ := http.NewRequest("GET", ts.URL, nil)
4310 res, err := tr.RoundTrip(req)
4311 if err != nil {
4312 t.Fatal(err)
4313 }
4314 res.Body.Close()
4315 <-donec
4316 }
4317
4318
4319
4320
4321
4322
4323 func TestServerWriteDoesNotRetainBufferAfterServerClose(t *testing.T) {
4324 synctest.Test(t, testServerWriteDoesNotRetainBufferAfterServerClose)
4325 }
4326 func testServerWriteDoesNotRetainBufferAfterServerClose(t *testing.T) {
4327 donec := make(chan struct{}, 1)
4328 ts := newTestServer(t, func(w http.ResponseWriter, r *http.Request) {
4329 donec <- struct{}{}
4330 defer close(donec)
4331 buf := make([]byte, 1<<20)
4332 var i byte
4333 for {
4334 i++
4335 _, err := w.Write(buf)
4336 for j := range buf {
4337 buf[j] = byte(i)
4338 }
4339 if err != nil {
4340 return
4341 }
4342 }
4343 })
4344
4345 tr := &http.Transport{
4346 TLSClientConfig: tlsConfigInsecure,
4347 Protocols: protocols("h2"),
4348 }
4349 defer tr.CloseIdleConnections()
4350
4351 req, _ := http.NewRequest("GET", ts.URL, nil)
4352 res, err := tr.RoundTrip(req)
4353 if err != nil {
4354 t.Fatal(err)
4355 }
4356 defer res.Body.Close()
4357 <-donec
4358 ts.Config.Close()
4359 <-donec
4360 }
4361
4362 func TestServerMaxHandlerGoroutines(t *testing.T) { synctest.Test(t, testServerMaxHandlerGoroutines) }
4363 func testServerMaxHandlerGoroutines(t *testing.T) {
4364 const maxHandlers = 10
4365 handlerc := make(chan chan bool)
4366 donec := make(chan struct{})
4367 defer close(donec)
4368 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4369 stopc := make(chan bool, 1)
4370 select {
4371 case handlerc <- stopc:
4372 case <-donec:
4373 }
4374 select {
4375 case shouldPanic := <-stopc:
4376 if shouldPanic {
4377 panic(http.ErrAbortHandler)
4378 }
4379 case <-donec:
4380 }
4381 }, func(h2 *http.HTTP2Config) {
4382 h2.MaxConcurrentStreams = maxHandlers
4383 })
4384 defer st.Close()
4385
4386 st.greet()
4387
4388
4389
4390 var stops []chan bool
4391 streamID := uint32(1)
4392 for range maxHandlers {
4393 st.writeHeaders(HeadersFrameParam{
4394 StreamID: streamID,
4395 BlockFragment: st.encodeHeader(),
4396 EndStream: true,
4397 EndHeaders: true,
4398 })
4399 stops = append(stops, <-handlerc)
4400 st.fr.WriteRSTStream(streamID, ErrCodeCancel)
4401 streamID += 2
4402 }
4403
4404
4405 st.writeHeaders(HeadersFrameParam{
4406 StreamID: streamID,
4407 BlockFragment: st.encodeHeader(),
4408 EndStream: true,
4409 EndHeaders: true,
4410 })
4411 st.fr.WriteRSTStream(streamID, ErrCodeCancel)
4412 streamID += 2
4413
4414
4415 for range 2 {
4416 st.writeHeaders(HeadersFrameParam{
4417 StreamID: streamID,
4418 BlockFragment: st.encodeHeader(),
4419 EndStream: true,
4420 EndHeaders: true,
4421 })
4422 streamID += 2
4423 }
4424
4425
4426
4427 select {
4428 case <-handlerc:
4429 t.Errorf("handler unexpectedly started while maxHandlers are already running")
4430 case <-time.After(1 * time.Millisecond):
4431 }
4432
4433
4434
4435 stops[0] <- false
4436 stops[1] <- true
4437 stops = stops[2:]
4438 stops = append(stops, <-handlerc)
4439 stops = append(stops, <-handlerc)
4440
4441
4442
4443 for range 5 * maxHandlers {
4444 st.writeHeaders(HeadersFrameParam{
4445 StreamID: streamID,
4446 BlockFragment: st.encodeHeader(),
4447 EndStream: true,
4448 EndHeaders: true,
4449 })
4450 st.fr.WriteRSTStream(streamID, ErrCodeCancel)
4451 streamID += 2
4452 }
4453 fr := readFrame[*GoAwayFrame](t, st)
4454 if fr.ErrCode != ErrCodeEnhanceYourCalm {
4455 t.Errorf("err code = %v; want %v", fr.ErrCode, ErrCodeEnhanceYourCalm)
4456 }
4457
4458 for _, s := range stops {
4459 close(s)
4460 }
4461 }
4462
4463 func TestServerContinuationFlood(t *testing.T) { synctest.Test(t, testServerContinuationFlood) }
4464 func testServerContinuationFlood(t *testing.T) {
4465 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4466 fmt.Println(r.Header)
4467 }, func(s *http.Server) {
4468 s.MaxHeaderBytes = 4096
4469 })
4470 defer st.Close()
4471
4472 st.greet()
4473
4474 st.writeHeaders(HeadersFrameParam{
4475 StreamID: 1,
4476 BlockFragment: st.encodeHeader(),
4477 EndStream: true,
4478 })
4479 for i := range 1000 {
4480 st.fr.WriteContinuation(1, false, st.encodeHeaderRaw(
4481 fmt.Sprintf("x-%v", i), "1234567890",
4482 ))
4483 }
4484 st.fr.WriteContinuation(1, true, st.encodeHeaderRaw(
4485 "x-last-header", "1",
4486 ))
4487
4488 for {
4489 f := st.readFrame()
4490 if f == nil {
4491 break
4492 }
4493 switch f := f.(type) {
4494 case *HeadersFrame:
4495 t.Fatalf("received HEADERS frame; want GOAWAY and a closed connection")
4496 case *GoAwayFrame:
4497
4498
4499
4500 if got, want := f.LastStreamID, uint32(1); got != want {
4501 t.Errorf("received GOAWAY with LastStreamId %v, want %v", got, want)
4502 }
4503
4504 }
4505 }
4506
4507
4508
4509
4510
4511
4512
4513
4514 }
4515
4516 func TestServerContinuationAfterInvalidHeader(t *testing.T) {
4517 synctest.Test(t, testServerContinuationAfterInvalidHeader)
4518 }
4519 func testServerContinuationAfterInvalidHeader(t *testing.T) {
4520 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4521 fmt.Println(r.Header)
4522 })
4523 defer st.Close()
4524
4525 st.greet()
4526
4527 st.writeHeaders(HeadersFrameParam{
4528 StreamID: 1,
4529 BlockFragment: st.encodeHeader(),
4530 EndStream: true,
4531 })
4532 st.fr.WriteContinuation(1, false, st.encodeHeaderRaw(
4533 "x-invalid-header", "\x00",
4534 ))
4535 st.fr.WriteContinuation(1, true, st.encodeHeaderRaw(
4536 "x-valid-header", "1",
4537 ))
4538
4539 var sawGoAway bool
4540 for {
4541 f := st.readFrame()
4542 if f == nil {
4543 break
4544 }
4545 switch f.(type) {
4546 case *GoAwayFrame:
4547 sawGoAway = true
4548 case *HeadersFrame:
4549 t.Fatalf("received HEADERS frame; want GOAWAY")
4550 }
4551 }
4552 if !sawGoAway {
4553 t.Errorf("connection closed with no GOAWAY frame; want one")
4554 }
4555 }
4556
4557
4558 func TestServerRequestCancelOnError(t *testing.T) { synctest.Test(t, testServerRequestCancelOnError) }
4559 func testServerRequestCancelOnError(t *testing.T) {
4560 recvc := make(chan struct{})
4561 donec := make(chan struct{})
4562 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4563 close(recvc)
4564 <-r.Context().Done()
4565 close(donec)
4566 })
4567 defer st.Close()
4568
4569 st.greet()
4570
4571
4572 st.writeHeaders(HeadersFrameParam{
4573 StreamID: 1,
4574 BlockFragment: st.encodeHeader(),
4575 EndStream: true,
4576 EndHeaders: true,
4577 })
4578 <-recvc
4579
4580
4581
4582
4583 st.writeHeaders(HeadersFrameParam{
4584 StreamID: 1,
4585 BlockFragment: st.encodeHeader(),
4586 EndStream: true,
4587 EndHeaders: true,
4588 })
4589 <-donec
4590 }
4591
4592 func TestServerSetReadWriteDeadlineRace(t *testing.T) {
4593 synctest.Test(t, testServerSetReadWriteDeadlineRace)
4594 }
4595 func testServerSetReadWriteDeadlineRace(t *testing.T) {
4596 ts := newTestServer(t, func(w http.ResponseWriter, r *http.Request) {
4597 ctl := http.NewResponseController(w)
4598 ctl.SetReadDeadline(time.Now().Add(3600 * time.Second))
4599 ctl.SetWriteDeadline(time.Now().Add(3600 * time.Second))
4600 })
4601 resp, err := ts.Client().Get(ts.URL)
4602 if err != nil {
4603 t.Fatal(err)
4604 }
4605 resp.Body.Close()
4606 }
4607
4608 func TestServerWriteByteTimeout(t *testing.T) { synctest.Test(t, testServerWriteByteTimeout) }
4609 func testServerWriteByteTimeout(t *testing.T) {
4610 const timeout = 1 * time.Second
4611 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4612 w.Write(make([]byte, 100))
4613 }, func(s *http.Server) {
4614
4615
4616
4617 s.Protocols = protocols("h2c")
4618 }, func(h2 *http.HTTP2Config) {
4619 h2.WriteByteTimeout = timeout
4620 })
4621 st.greet()
4622
4623 st.cc.(*synctestNetConn).SetReadBufferSize(1)
4624 st.writeHeaders(HeadersFrameParam{
4625 StreamID: 1,
4626 BlockFragment: st.encodeHeader(),
4627 EndStream: true,
4628 EndHeaders: true,
4629 })
4630
4631
4632 for i := range 10 {
4633 st.advance(timeout - 1)
4634 if n, err := st.cc.Read(make([]byte, 1)); n != 1 || err != nil {
4635 t.Fatalf("read %v: %v, %v; want 1, nil", i, n, err)
4636 }
4637 }
4638
4639
4640
4641 st.advance(1 * time.Second)
4642 st.advance(1 * time.Second)
4643 st.wantClosed()
4644 }
4645
4646 func TestServerPingSent(t *testing.T) { synctest.Test(t, testServerPingSent) }
4647 func testServerPingSent(t *testing.T) {
4648 const sendPingTimeout = 15 * time.Second
4649 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4650 }, func(h2 *http.HTTP2Config) {
4651 h2.SendPingTimeout = sendPingTimeout
4652 })
4653 st.greet()
4654
4655 st.wantIdle()
4656
4657 st.advance(sendPingTimeout)
4658 _ = readFrame[*PingFrame](t, st)
4659 st.wantIdle()
4660
4661 st.advance(14 * time.Second)
4662 st.wantIdle()
4663 st.advance(1 * time.Second)
4664 st.wantClosed()
4665 }
4666
4667 func TestServerPingResponded(t *testing.T) { synctest.Test(t, testServerPingResponded) }
4668 func testServerPingResponded(t *testing.T) {
4669 const sendPingTimeout = 15 * time.Second
4670 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4671 }, func(h2 *http.HTTP2Config) {
4672 h2.SendPingTimeout = sendPingTimeout
4673 })
4674 st.greet()
4675
4676 st.wantIdle()
4677
4678 st.advance(sendPingTimeout)
4679 pf := readFrame[*PingFrame](t, st)
4680 st.wantIdle()
4681
4682 st.advance(14 * time.Second)
4683 st.wantIdle()
4684
4685 st.writePing(true, pf.Data)
4686
4687 st.advance(2 * time.Second)
4688 st.wantIdle()
4689 }
4690
4691
4692
4693
4694
4695 func TestServerSendDataAfterRequestBodyClose(t *testing.T) {
4696 synctest.Test(t, testServerSendDataAfterRequestBodyClose)
4697 }
4698 func testServerSendDataAfterRequestBodyClose(t *testing.T) {
4699 st := newServerTester(t, nil)
4700 st.greet()
4701
4702 st.writeHeaders(HeadersFrameParam{
4703 StreamID: 1,
4704 BlockFragment: st.encodeHeader(),
4705 EndStream: false,
4706 EndHeaders: true,
4707 })
4708
4709
4710 call := st.nextHandlerCall()
4711 call.do(func(w http.ResponseWriter, req *http.Request) {
4712 w.Write([]byte("one"))
4713 http.NewResponseController(w).Flush()
4714 })
4715 st.wantFrameType(FrameHeaders)
4716 st.wantData(wantData{
4717 streamID: 1,
4718 endStream: false,
4719 data: []byte("one"),
4720 })
4721 st.wantIdle()
4722
4723
4724
4725 call.do(func(w http.ResponseWriter, req *http.Request) {
4726 req.Body.Close()
4727 })
4728 st.wantIdle()
4729
4730
4731 st.writeData(1, false, []byte("client-sent data"))
4732 st.wantIdle()
4733
4734
4735
4736 call.do(func(w http.ResponseWriter, req *http.Request) {
4737 w.Write([]byte("two"))
4738 http.NewResponseController(w).Flush()
4739 })
4740 st.wantData(wantData{
4741 streamID: 1,
4742 endStream: false,
4743 data: []byte("two"),
4744 })
4745 st.wantIdle()
4746 }
4747
4748 func TestServerSettingNoRFC7540Priorities(t *testing.T) {
4749 synctest.Test(t, testServerSettingNoRFC7540Priorities)
4750 }
4751 func testServerSettingNoRFC7540Priorities(t *testing.T) {
4752 const wantNoRFC7540Setting = true
4753 st := newServerTester(t, nil)
4754 defer st.Close()
4755
4756 var gotNoRFC7540Setting bool
4757 st.greetAndCheckSettings(func(s Setting) error {
4758 if s.ID != SettingNoRFC7540Priorities {
4759 return nil
4760 }
4761 gotNoRFC7540Setting = s.Val == 1
4762 return nil
4763 })
4764 if wantNoRFC7540Setting != gotNoRFC7540Setting {
4765 t.Errorf("want SETTINGS_NO_RFC7540_PRIORITIES to be %v, got %v", wantNoRFC7540Setting, gotNoRFC7540Setting)
4766 }
4767 }
4768
4769 func TestServerSettingNoRFC7540PrioritiesInvalid(t *testing.T) {
4770 synctest.Test(t, testServerSettingNoRFC7540PrioritiesInvalid)
4771 }
4772 func testServerSettingNoRFC7540PrioritiesInvalid(t *testing.T) {
4773 st := newServerTester(t, nil)
4774 defer st.Close()
4775
4776 st.writePreface()
4777 st.writeSettings(Setting{ID: SettingNoRFC7540Priorities, Val: 2})
4778 synctest.Wait()
4779 st.readFrame()
4780 st.readFrame()
4781 st.wantGoAway(0, ErrCodeProtocol)
4782 }
4783
4784
4785
4786 func TestServerRFC9218PrioritySmallPayload(t *testing.T) {
4787 synctest.Test(t, testServerRFC9218PrioritySmallPayload)
4788 }
4789 func testServerRFC9218PrioritySmallPayload(t *testing.T) {
4790 endTest := false
4791 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4792 for !endTest {
4793 w.Write([]byte("a"))
4794 if f, ok := w.(http.Flusher); ok {
4795 f.Flush()
4796 }
4797 }
4798 }, func(s *http.Server) {
4799 s.Protocols = protocols("h2c")
4800 })
4801 st.greet()
4802 if syncConn, ok := st.cc.(*synctestNetConn); ok {
4803 syncConn.SetReadBufferSize(1)
4804 } else {
4805 t.Fatal("Server connection is not synctestNetConn")
4806 }
4807 defer st.Close()
4808 defer func() { endTest = true }()
4809
4810
4811
4812
4813
4814
4815 for i := 1; i <= 19; i += 2 {
4816 urgency := uint8(0)
4817 if i > 10 {
4818 urgency = 7
4819 }
4820 st.writeHeaders(HeadersFrameParam{
4821 StreamID: uint32(i),
4822 BlockFragment: st.encodeHeader("priority", fmt.Sprintf("u=%d", urgency)),
4823 EndStream: true,
4824 EndHeaders: true,
4825 })
4826 synctest.Wait()
4827 }
4828
4829
4830
4831 streamWriteCount := make(map[uint32]int)
4832 totalWriteCount := 10000
4833 for range totalWriteCount {
4834 f := st.readFrame()
4835 if f == nil {
4836 break
4837 }
4838 streamWriteCount[f.Header().StreamID] += 1
4839 }
4840 for streamID, writeCount := range streamWriteCount {
4841 expectedWriteCount := totalWriteCount / len(streamWriteCount)
4842 errorMargin := expectedWriteCount / 100
4843 if writeCount >= expectedWriteCount+errorMargin || writeCount <= expectedWriteCount-errorMargin {
4844 t.Errorf("Expected stream %v to receive %v±%v writes, got %v", streamID, expectedWriteCount, errorMargin, writeCount)
4845 }
4846 }
4847 }
4848
4849 func TestServerRFC9218Priority(t *testing.T) {
4850 synctest.Test(t, testServerRFC9218Priority)
4851 }
4852 func testServerRFC9218Priority(t *testing.T) {
4853 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4854 w.Write(slices.Repeat([]byte("a"), 16<<20))
4855 if f, ok := w.(http.Flusher); ok {
4856 f.Flush()
4857 }
4858 }, func(s *http.Server) {
4859 s.Protocols = protocols("h2c")
4860 })
4861 defer st.Close()
4862 st.greet()
4863 if syncConn, ok := st.cc.(*synctestNetConn); ok {
4864 syncConn.SetReadBufferSize(1)
4865 } else {
4866 t.Fatal("Server connection is not synctestNetConn")
4867 }
4868 st.writeWindowUpdate(0, 1<<30)
4869 synctest.Wait()
4870
4871
4872
4873 for i := range 8 {
4874 streamID := uint32(i*2 + 1)
4875 urgency := 7 - i
4876 st.writeHeaders(HeadersFrameParam{
4877 StreamID: streamID,
4878 BlockFragment: st.encodeHeader("priority", fmt.Sprintf("u=%d", urgency)),
4879 EndStream: true,
4880 EndHeaders: true,
4881 })
4882 }
4883 synctest.Wait()
4884
4885
4886
4887 lastFrame := make(map[uint32]int)
4888 for i := 0; ; i++ {
4889 f := st.readFrame()
4890 if f == nil {
4891 break
4892 }
4893 lastFrame[f.Header().StreamID] = i
4894 }
4895 for i := range 7 {
4896 streamID := uint32(i*2 + 1)
4897 nextStreamID := streamID + 2
4898 if lastFrame[streamID] < lastFrame[nextStreamID] {
4899 t.Errorf("stream %d finished before stream %d unexpectedly", streamID, nextStreamID)
4900 }
4901 }
4902 }
4903
4904 func TestServerRFC9218PriorityIgnoredWhenProxied(t *testing.T) {
4905 synctest.Test(t, testServerRFC9218PriorityIgnoredWhenProxied)
4906 }
4907 func testServerRFC9218PriorityIgnoredWhenProxied(t *testing.T) {
4908 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4909 w.Write(slices.Repeat([]byte("a"), 16<<20))
4910 if f, ok := w.(http.Flusher); ok {
4911 f.Flush()
4912 }
4913 }, func(s *http.Server) {
4914 s.Protocols = protocols("h2c")
4915 })
4916 defer st.Close()
4917 st.greet()
4918 if syncConn, ok := st.cc.(*synctestNetConn); ok {
4919 syncConn.SetReadBufferSize(1)
4920 } else {
4921 t.Fatal("Server connection is not synctestNetConn")
4922 }
4923 st.writeWindowUpdate(0, 1<<30)
4924 synctest.Wait()
4925
4926
4927
4928
4929 for i := range 8 {
4930 streamID := uint32(i*2 + 1)
4931 urgency := 7 - i
4932 st.writeHeaders(HeadersFrameParam{
4933 StreamID: streamID,
4934 BlockFragment: st.encodeHeader("priority", fmt.Sprintf("u=%d", urgency), "via", "a proxy"),
4935 EndStream: true,
4936 EndHeaders: true,
4937 })
4938 }
4939 synctest.Wait()
4940 var streamFrameOrder []uint32
4941 for f := st.readFrame(); f != nil; f = st.readFrame() {
4942 streamFrameOrder = append(streamFrameOrder, f.Header().StreamID)
4943 }
4944
4945
4946
4947 half := streamFrameOrder[len(streamFrameOrder)/4 : len(streamFrameOrder)*3/4]
4948 if !slices.Equal(slices.Compact(half), half) {
4949 t.Errorf("want stream to be processed in round-robin manner when proxied, got: %v", streamFrameOrder)
4950 }
4951 }
4952
4953 func TestServerRFC9218PriorityAware(t *testing.T) {
4954 synctest.Test(t, testServerRFC9218PriorityAware)
4955 }
4956 func testServerRFC9218PriorityAware(t *testing.T) {
4957 st := newServerTester(t, func(w http.ResponseWriter, r *http.Request) {
4958 w.Write(slices.Repeat([]byte("a"), 16<<20))
4959 if f, ok := w.(http.Flusher); ok {
4960 f.Flush()
4961 }
4962 }, func(s *http.Server) {
4963 s.Protocols = protocols("h2c")
4964 })
4965 defer st.Close()
4966 st.greet()
4967 if syncConn, ok := st.cc.(*synctestNetConn); ok {
4968 syncConn.SetReadBufferSize(1)
4969 } else {
4970 t.Fatal("Server connection is not synctestNetConn")
4971 }
4972 st.writeWindowUpdate(0, 1<<30)
4973 synctest.Wait()
4974
4975
4976
4977 streamCount := 10
4978 for i := range streamCount {
4979 streamID := uint32(i*2 + 1)
4980 st.writeHeaders(HeadersFrameParam{
4981 StreamID: streamID,
4982 BlockFragment: st.encodeHeader(),
4983 EndStream: true,
4984 EndHeaders: true,
4985 })
4986 }
4987 synctest.Wait()
4988 var streamFrameOrder []uint32
4989 for f := st.readFrame(); f != nil; f = st.readFrame() {
4990 streamFrameOrder = append(streamFrameOrder, f.Header().StreamID)
4991 }
4992
4993
4994
4995 half := streamFrameOrder[len(streamFrameOrder)/4 : len(streamFrameOrder)*3/4]
4996 if !slices.Equal(slices.Compact(half), half) {
4997 t.Errorf("want stream to be processed in round-robin manner when unaware of priority, got: %v", streamFrameOrder)
4998 }
4999
5000
5001
5002
5003 st.writePriorityUpdate(1, "")
5004 synctest.Wait()
5005
5006
5007
5008
5009
5010 streamFrameOrder = []uint32{}
5011 for i := range streamCount {
5012 i += streamCount
5013 streamID := uint32(i*2 + 1)
5014 st.writeHeaders(HeadersFrameParam{
5015 StreamID: streamID,
5016 BlockFragment: st.encodeHeader(),
5017 EndStream: true,
5018 EndHeaders: true,
5019 })
5020 }
5021 for f := st.readFrame(); f != nil; f = st.readFrame() {
5022 streamFrameOrder = append(streamFrameOrder, f.Header().StreamID)
5023 }
5024 if !slices.Equal(slices.Compact(half), half) {
5025 t.Errorf("want stream to be processed one-by-one to completion when aware of priority, got: %v", streamFrameOrder)
5026 }
5027 }
5028
5029 func TestServerInvalidPathHeader(t *testing.T) {
5030 synctest.Test(t, testServerInvalidPathHeader)
5031 }
5032 func testServerInvalidPathHeader(t *testing.T) {
5033 for _, path := range []string{
5034 "",
5035 "\x00",
5036 "https://example.com/",
5037 } {
5038 testServerRejectsStream(t, ErrCodeProtocol, func(st *serverTester) {
5039 st.fr.AllowIllegalWrites = true
5040 st.writeHeaders(HeadersFrameParam{
5041 StreamID: 1,
5042 BlockFragment: st.encodeHeader(
5043 ":path", path,
5044 ),
5045 EndStream: true,
5046 EndHeaders: true,
5047 })
5048 })
5049 }
5050 }
5051
5052 func TestServerPathInitialSlashes(t *testing.T) {
5053 synctest.Test(t, testServerPathInitialSlashes)
5054 }
5055 func testServerPathInitialSlashes(t *testing.T) {
5056 st := newServerTester(t, nil)
5057 st.greet()
5058
5059
5060
5061 const path = "//narf.com/path"
5062 st.writeHeaders(HeadersFrameParam{
5063 StreamID: 1,
5064 BlockFragment: st.encodeHeader(
5065 ":path", path,
5066 ),
5067 EndStream: true,
5068 EndHeaders: true,
5069 })
5070
5071 call := st.nextHandlerCall()
5072 if got, want := call.req.URL.Host, ""; got != want {
5073 t.Errorf("got req.URL.Host %q, want %q", got, want)
5074 }
5075 if got, want := call.req.URL.Path, path; got != want {
5076 t.Errorf("got req.URL.Path %q, want %q", got, want)
5077 }
5078 }
5079
5080
5081
5082
5083
5084 func TestServerSettingsFlowControlUpdateBeyondLimit(t *testing.T) {
5085 synctest.Test(t, testServerSettingsFlowControlUpdateBeyondLimit)
5086 }
5087 func testServerSettingsFlowControlUpdateBeyondLimit(t *testing.T) {
5088 st := newServerTester(t, nil)
5089 st.greet()
5090
5091 st.writeHeaders(HeadersFrameParam{
5092 StreamID: 1,
5093 BlockFragment: st.encodeHeader(":method", "POST"),
5094 EndStream: false,
5095 EndHeaders: true,
5096 })
5097
5098
5099 const windowIncrease = 1000
5100 st.writeWindowUpdate(1, windowIncrease)
5101 st.wantIdle()
5102
5103
5104 const maxWindowSize = (1 << 31) - 1
5105 const maxInitialWindowSize = maxWindowSize - windowIncrease
5106 st.writeSettings(Setting{SettingInitialWindowSize, maxInitialWindowSize + 1})
5107 st.wantGoAway(1, ErrCodeFlowControl)
5108 }
5109
5110
5111
5112 func TestServerSettingsFlowControlUpdateWithinLimit(t *testing.T) {
5113 synctest.Test(t, testServerSettingsFlowControlUpdateWithinLimit)
5114 }
5115 func testServerSettingsFlowControlUpdateWithinLimit(t *testing.T) {
5116 st := newServerTester(t, nil)
5117 st.greet()
5118
5119 st.writeHeaders(HeadersFrameParam{
5120 StreamID: 1,
5121 BlockFragment: st.encodeHeader(":method", "POST"),
5122 EndStream: false,
5123 EndHeaders: true,
5124 })
5125
5126
5127 const windowIncrease = 1000
5128 st.writeWindowUpdate(1, windowIncrease)
5129 st.wantIdle()
5130
5131
5132 const maxWindowSize = (1 << 31) - 1
5133 const maxInitialWindowSize = maxWindowSize - windowIncrease
5134 st.writeSettings(Setting{SettingInitialWindowSize, maxInitialWindowSize})
5135 st.wantSettingsAck()
5136 st.wantIdle()
5137 }
5138
5139 func TestConsistentConstants(t *testing.T) {
5140 if h1, h2 := http.DefaultMaxHeaderBytes, http2.DefaultMaxHeaderBytes; h1 != h2 {
5141 t.Errorf("DefaultMaxHeaderBytes: http (%v) != http2 (%v)", h1, h2)
5142 }
5143 if h1, h2 := http.TimeFormat, http2.TimeFormat; h1 != h2 {
5144 t.Errorf("TimeFormat: http (%v) != http2 (%v)", h1, h2)
5145 }
5146 }
5147
5148 var (
5149 testServerTLSConfig *tls.Config
5150 testClientTLSConfig *tls.Config
5151 )
5152
5153 func init() {
5154 cert, err := tls.X509KeyPair(testcert.LocalhostCert, testcert.LocalhostKey)
5155 if err != nil {
5156 panic(err)
5157 }
5158 testServerTLSConfig = &tls.Config{
5159 Certificates: []tls.Certificate{cert},
5160 NextProtos: []string{"h2"},
5161 }
5162
5163 x509Cert, err := x509.ParseCertificate(cert.Certificate[0])
5164 if err != nil {
5165 panic(err)
5166 }
5167 certpool := x509.NewCertPool()
5168 certpool.AddCert(x509Cert)
5169 testClientTLSConfig = &tls.Config{
5170 InsecureSkipVerify: true,
5171 RootCAs: certpool,
5172 NextProtos: []string{"h2"},
5173 }
5174 }
5175
5176 func protocols(protos ...string) *http.Protocols {
5177 p := new(http.Protocols)
5178 for _, s := range protos {
5179 switch s {
5180 case "h1":
5181 p.SetHTTP1(true)
5182 case "h2":
5183 p.SetHTTP2(true)
5184 case "h2c":
5185 p.SetUnencryptedHTTP2(true)
5186 default:
5187 panic("unknown protocol: " + s)
5188 }
5189 }
5190 return p
5191 }
5192
5193
5194 func transportFromH1Transport(tr *http.Transport) any
5195
View as plain text