Source file src/net/http/internal/http3/stream_test.go

     1  // Copyright 2025 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     4  
     5  package http3
     6  
     7  import (
     8  	"bytes"
     9  	"errors"
    10  	"io"
    11  	"testing"
    12  )
    13  
    14  func TestStreamReadVarint(t *testing.T) {
    15  	st1, st2 := newStreamPair(t)
    16  	for _, b := range [][]byte{
    17  		{0x00},
    18  		{0x3f},
    19  		{0x40, 0x00},
    20  		{0x7f, 0xff},
    21  		{0x80, 0x00, 0x00, 0x00},
    22  		{0xbf, 0xff, 0xff, 0xff},
    23  		{0xc0, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00},
    24  		{0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff},
    25  		// Example cases from https://www.rfc-editor.org/rfc/rfc9000.html#section-a.1
    26  		{0xc2, 0x19, 0x7c, 0x5e, 0xff, 0x14, 0xe8, 0x8c},
    27  		{0x9d, 0x7f, 0x3e, 0x7d},
    28  		{0x7b, 0xbd},
    29  		{0x25},
    30  		{0x40, 0x25},
    31  	} {
    32  		trailer := []byte{0xde, 0xad, 0xbe, 0xef}
    33  		st1.Write(b)
    34  		st1.Write(trailer)
    35  		if err := st1.Flush(); err != nil {
    36  			t.Fatal(err)
    37  		}
    38  		got, err := st2.readVarint()
    39  		if err != nil {
    40  			t.Fatalf("st.readVarint() = %v", err)
    41  		}
    42  		want, _ := consumeVarintInt64(b)
    43  		if got != want {
    44  			t.Fatalf("st.readVarint() = %v, want %v", got, want)
    45  		}
    46  		gotTrailer := make([]byte, len(trailer))
    47  		if _, err := io.ReadFull(st2, gotTrailer); err != nil {
    48  			t.Fatal(err)
    49  		}
    50  		if !bytes.Equal(gotTrailer, trailer) {
    51  			t.Fatalf("after st.readVarint, read %x, want %x", gotTrailer, trailer)
    52  		}
    53  	}
    54  }
    55  
    56  func TestStreamWriteVarint(t *testing.T) {
    57  	st1, st2 := newStreamPair(t)
    58  	for _, v := range []int64{
    59  		0,
    60  		63,
    61  		16383,
    62  		1073741823,
    63  		4611686018427387903,
    64  		// Example cases from https://www.rfc-editor.org/rfc/rfc9000.html#section-a.1
    65  		151288809941952652,
    66  		494878333,
    67  		15293,
    68  		37,
    69  	} {
    70  		trailer := []byte{0xde, 0xad, 0xbe, 0xef}
    71  		st1.writeVarint(v)
    72  		st1.Write(trailer)
    73  		if err := st1.Flush(); err != nil {
    74  			t.Fatal(err)
    75  		}
    76  
    77  		want := appendVarint(nil, uint64(v))
    78  		want = append(want, trailer...)
    79  
    80  		got := make([]byte, len(want))
    81  		if _, err := io.ReadFull(st2, got); err != nil {
    82  			t.Fatal(err)
    83  		}
    84  
    85  		if !bytes.Equal(got, want) {
    86  			t.Errorf("AppendVarint(nil, %v) = %x, want %x", v, got, want)
    87  		}
    88  	}
    89  }
    90  
    91  func TestStreamReadFrames(t *testing.T) {
    92  	st1, st2 := newStreamPair(t)
    93  	for _, frame := range []struct {
    94  		ftype frameType
    95  		data  []byte
    96  	}{{
    97  		ftype: 1,
    98  		data:  []byte("hello"),
    99  	}, {
   100  		ftype: 2,
   101  		data:  []byte{},
   102  	}, {
   103  		ftype: 3,
   104  		data:  []byte("goodbye"),
   105  	}} {
   106  		st1.writeVarint(int64(frame.ftype))
   107  		st1.writeVarint(int64(len(frame.data)))
   108  		st1.Write(frame.data)
   109  		if err := st1.Flush(); err != nil {
   110  			t.Fatal(err)
   111  		}
   112  
   113  		if gotFrameType, err := st2.readFrameHeader(); err != nil || gotFrameType != frame.ftype {
   114  			t.Fatalf("st.readFrameHeader() = %v, %v; want %v, nil", gotFrameType, err, frame.ftype)
   115  		}
   116  		if gotData, err := st2.readFrameData(); err != nil || !bytes.Equal(gotData, frame.data) {
   117  			t.Fatalf("st.readFrameData() = %x, %v; want %x, nil", gotData, err, frame.data)
   118  		}
   119  		if err := st2.endFrame(); err != nil {
   120  			t.Fatalf("st.endFrame() = %v; want nil", err)
   121  		}
   122  	}
   123  }
   124  
   125  func TestStreamReadFrameUnderflow(t *testing.T) {
   126  	const size = 4
   127  	st1, st2 := newStreamPair(t)
   128  	st1.writeVarint(0)            // type
   129  	st1.writeVarint(size)         // size
   130  	st1.Write(make([]byte, size)) // data
   131  	if err := st1.Flush(); err != nil {
   132  		t.Fatal(err)
   133  	}
   134  
   135  	if _, err := st2.readFrameHeader(); err != nil {
   136  		t.Fatalf("st.readFrameHeader() = %v", err)
   137  	}
   138  	if _, err := io.ReadFull(st2, make([]byte, size-1)); err != nil {
   139  		t.Fatalf("st.Read() = %v", err)
   140  	}
   141  	// We have not consumed the full frame: Error.
   142  	if err := st2.endFrame(); !errors.Is(err, errH3FrameError) {
   143  		t.Fatalf("st.endFrame before end: %v, want errH3FrameError", err)
   144  	}
   145  }
   146  
   147  func TestStreamReadFrameWithoutEnd(t *testing.T) {
   148  	const size = 4
   149  	st1, st2 := newStreamPair(t)
   150  	st1.writeVarint(0)            // type
   151  	st1.writeVarint(size)         // size
   152  	st1.Write(make([]byte, size)) // data
   153  	if err := st1.Flush(); err != nil {
   154  		t.Fatal(err)
   155  	}
   156  
   157  	if _, err := st2.readFrameHeader(); err != nil {
   158  		t.Fatalf("st.readFrameHeader() = %v", err)
   159  	}
   160  	if _, err := st2.readFrameHeader(); err == nil {
   161  		t.Fatalf("st.readFrameHeader before st.endFrame for prior frame: success, want error")
   162  	}
   163  }
   164  
   165  func TestStreamReadFrameOverflow(t *testing.T) {
   166  	const size = 4
   167  	st1, st2 := newStreamPair(t)
   168  	st1.writeVarint(0)              // type
   169  	st1.writeVarint(size)           // size
   170  	st1.Write(make([]byte, size+1)) // data
   171  	if err := st1.Flush(); err != nil {
   172  		t.Fatal(err)
   173  	}
   174  
   175  	if _, err := st2.readFrameHeader(); err != nil {
   176  		t.Fatalf("st.readFrameHeader() = %v", err)
   177  	}
   178  	if _, err := io.ReadFull(st2, make([]byte, size+1)); !errors.Is(err, errH3FrameError) {
   179  		t.Fatalf("st.Read past end of frame: %v, want errH3FrameError", err)
   180  	}
   181  }
   182  
   183  func TestStreamReadFrameHeaderPartial(t *testing.T) {
   184  	var frame []byte
   185  	frame = appendVarint(frame, 1000) // type
   186  	frame = appendVarint(frame, 2000) // size
   187  
   188  	for i := 1; i < len(frame)-1; i++ {
   189  		st1, st2 := newStreamPair(t)
   190  		st1.Write(frame[:i])
   191  		if err := st1.Flush(); err != nil {
   192  			t.Fatal(err)
   193  		}
   194  		st1.CloseWrite()
   195  
   196  		if _, err := st2.readFrameHeader(); err == nil {
   197  			t.Fatalf("%v/%v bytes of frame available: st.readFrameHeader() succeeded; want error", i, len(frame))
   198  		}
   199  	}
   200  }
   201  
   202  func TestStreamReadFrameDataPartial(t *testing.T) {
   203  	st1, st2 := newStreamPair(t)
   204  	st1.writeVarint(1)          // type
   205  	st1.writeVarint(100)        // size
   206  	st1.Write(make([]byte, 50)) // data
   207  	st1.CloseWrite()
   208  	if _, err := st2.readFrameHeader(); err != nil {
   209  		t.Fatalf("st.readFrameHeader() = %v", err)
   210  	}
   211  	if n, err := io.ReadAll(st2); err == nil {
   212  		t.Fatalf("io.ReadAll with partial frame = %v, nil; want error", n)
   213  	}
   214  }
   215  
   216  func TestStreamReadByteFrameDataPartial(t *testing.T) {
   217  	st1, st2 := newStreamPair(t)
   218  	st1.writeVarint(1)   // type
   219  	st1.writeVarint(100) // size
   220  	st1.CloseWrite()
   221  	if _, err := st2.readFrameHeader(); err != nil {
   222  		t.Fatalf("st.readFrameHeader() = %v", err)
   223  	}
   224  	if b, err := st2.ReadByte(); err == nil {
   225  		t.Fatalf("io.ReadAll with partial frame = %v, nil; want error", b)
   226  	}
   227  }
   228  
   229  func TestStreamReadFrameDataAtEOF(t *testing.T) {
   230  	const typ = 10
   231  	data := []byte("hello")
   232  	st1, st2 := newStreamPair(t)
   233  	st1.writeVarint(typ)              // type
   234  	st1.writeVarint(int64(len(data))) // size
   235  	if err := st1.Flush(); err != nil {
   236  		t.Fatal(err)
   237  	}
   238  	if got, err := st2.readFrameHeader(); err != nil || got != typ {
   239  		t.Fatalf("st.readFrameHeader() = %v, %v; want %v, nil", got, err, typ)
   240  	}
   241  
   242  	st1.Write(data)         // data
   243  	st1.CloseWrite() // end stream
   244  	got := make([]byte, len(data)+1)
   245  	if n, err := st2.Read(got); err != nil || n != len(data) || !bytes.Equal(got[:n], data) {
   246  		t.Fatalf("st.Read() = %v, %v (data=%x); want %v, nil (data=%x)", n, err, got[:n], len(data), data)
   247  	}
   248  }
   249  
   250  func TestStreamReadFrameData(t *testing.T) {
   251  	const typ = 10
   252  	data := []byte("hello")
   253  	st1, st2 := newStreamPair(t)
   254  	st1.writeVarint(typ)              // type
   255  	st1.writeVarint(int64(len(data))) // size
   256  	st1.Write(data)                   // data
   257  	if err := st1.Flush(); err != nil {
   258  		t.Fatal(err)
   259  	}
   260  
   261  	if got, err := st2.readFrameHeader(); err != nil || got != typ {
   262  		t.Fatalf("st.readFrameHeader() = %v, %v; want %v, nil", got, err, typ)
   263  	}
   264  	if got, err := st2.readFrameData(); err != nil || !bytes.Equal(got, data) {
   265  		t.Fatalf("st.readFrameData() = %x, %v; want %x, nil", got, err, data)
   266  	}
   267  }
   268  
   269  func TestStreamReadByte(t *testing.T) {
   270  	const stype = 1
   271  	const want = 42
   272  	st1, st2 := newStreamPair(t)
   273  	st1.writeVarint(stype)  // stream type
   274  	st1.writeVarint(1)      // size
   275  	st1.Write([]byte{want}) // data
   276  	if err := st1.Flush(); err != nil {
   277  		t.Fatal(err)
   278  	}
   279  
   280  	if got, err := st2.readFrameHeader(); err != nil || got != stype {
   281  		t.Fatalf("st.readFrameHeader() = %v, %v; want %v, nil", got, err, stype)
   282  	}
   283  	if got, err := st2.ReadByte(); err != nil || got != want {
   284  		t.Fatalf("st.ReadByte() = %v, %v; want %v, nil", got, err, want)
   285  	}
   286  	if got, err := st2.ReadByte(); err == nil {
   287  		t.Fatalf("reading past end of frame: st.ReadByte() = %v, %v; want error", got, err)
   288  	}
   289  }
   290  
   291  func TestStreamDiscardFrame(t *testing.T) {
   292  	const typ = 10
   293  	data := []byte("hello")
   294  	st1, st2 := newStreamPair(t)
   295  	st1.writeVarint(typ)              // type
   296  	st1.writeVarint(int64(len(data))) // size
   297  	st1.Write(data)                   // data
   298  	st1.CloseWrite()
   299  
   300  	if got, err := st2.readFrameHeader(); err != nil || got != typ {
   301  		t.Fatalf("st.readFrameHeader() = %v, %v; want %v, nil", got, err, typ)
   302  	}
   303  	if err := st2.discardFrame(); err != nil {
   304  		t.Fatalf("st.discardFrame() = %v", err)
   305  	}
   306  	if b, err := io.ReadAll(st2); err != nil || len(b) > 0 {
   307  		t.Fatalf("after discarding frame, read %x, %v; want EOF", b, err)
   308  	}
   309  }
   310  
   311  func newStreamPair(t testing.TB) (s1, s2 *stream) {
   312  	t.Helper()
   313  	q1, q2 := newQUICStreamPair(t)
   314  	return newStream(q1), newStream(q2)
   315  }
   316  

View as plain text