1
2
3
4
5 package http3
6
7 import (
8 "errors"
9 "io"
10 "net/http"
11 "net/http/httptrace"
12 "net/textproto"
13 "strconv"
14 "net/http/internal/ascii"
15 "sync"
16
17 "golang.org/x/net/http/httpguts"
18 "net/http/internal/httpcommon"
19 "golang.org/x/net/quic"
20 )
21
22 type roundTripState struct {
23 cc *clientConn
24 st *stream
25
26
27 onceCloseReqBody sync.Once
28 reqBody io.ReadCloser
29
30 reqBodyWriter bodyWriter
31
32
33 respBody io.ReadCloser
34
35 trace *httptrace.ClientTrace
36
37 errOnce sync.Once
38 err error
39 }
40
41
42
43 func (rt *roundTripState) abort(err error) error {
44 rt.errOnce.Do(func() {
45 rt.err = err
46
47 rt.cc.mu.Lock()
48 rt.cc.active--
49 rt.cc.mu.Unlock()
50 rt.cc.maybeCallStateHook()
51
52 switch e := err.(type) {
53 case *connectionError:
54 rt.cc.abort(e)
55 case *streamError:
56 rt.st.CloseRead(uint64(e.code))
57 rt.st.Reset(uint64(e.code))
58 default:
59 rt.st.CloseRead(uint64(errH3NoError))
60 rt.st.Reset(uint64(errH3NoError))
61 }
62 })
63 return rt.err
64 }
65
66
67 func (rt *roundTripState) closeReqBody() {
68 if rt.reqBody != nil {
69 rt.onceCloseReqBody.Do(func() {
70 rt.reqBody.Close()
71 })
72 }
73 }
74
75
76 func (rt *roundTripState) maybeCallGot1xxResponse(status int, h http.Header) error {
77 if rt.trace == nil || rt.trace.Got1xxResponse == nil {
78 return nil
79 }
80 return rt.trace.Got1xxResponse(status, textproto.MIMEHeader(h))
81 }
82
83 func (rt *roundTripState) maybeCallGot100Continue() {
84 if rt.trace == nil || rt.trace.Got100Continue == nil {
85 return
86 }
87 rt.trace.Got100Continue()
88 }
89
90 func (rt *roundTripState) maybeCallWait100Continue() {
91 if rt.trace == nil || rt.trace.Wait100Continue == nil {
92 return
93 }
94 rt.trace.Wait100Continue()
95 }
96
97
98 func (cc *clientConn) RoundTrip(req *http.Request) (_ *http.Response, err error) {
99 cc.mu.Lock()
100 if cc.reserved > 0 {
101 cc.reserved--
102 }
103 cc.active++
104 cc.mu.Unlock()
105
106
107 st, err := newConnStream(req.Context(), cc.qconn, streamTypeRequest)
108 if err != nil {
109 cc.mu.Lock()
110 cc.active--
111 cc.mu.Unlock()
112 cc.maybeCallStateHook()
113 return nil, err
114 }
115 rt := &roundTripState{
116 cc: cc,
117 st: st,
118 trace: httptrace.ContextClientTrace(req.Context()),
119 reqBody: req.Body,
120 }
121 if rt.reqBody == nil {
122 rt.reqBody = http.NoBody
123 }
124
125 var wg sync.WaitGroup
126 defer func() {
127 if err != nil {
128 err = rt.abort(err)
129
130
131
132
133
134
135
136
137
138
139
140
141
142 rt.closeReqBody()
143 wg.Wait()
144 }
145 }()
146
147
148 st.stream.SetReadContext(req.Context())
149 st.stream.SetWriteContext(req.Context())
150
151 addedGzip := httpcommon.IsRequestGzip(req.Method, req.Header, cc.tr.tr1.DisableCompression)
152 headers := cc.enc.encode(func(yield func(itype indexType, name, value string)) {
153 _, err = httpcommon.EncodeHeaders(req.Context(), httpcommon.EncodeHeadersParam{
154 Request: httpcommon.Request{
155 URL: req.URL,
156 Method: req.Method,
157 Host: req.Host,
158 Header: req.Header,
159 Trailer: req.Trailer,
160 ActualContentLength: actualContentLength(req),
161 },
162 AddGzipHeader: addedGzip,
163 PeerMaxHeaderListSize: 0,
164 DefaultUserAgent: "Go-http-client/3.0",
165 }, func(name, value string) {
166
167 yield(mayIndex, name, value)
168 })
169 })
170 if err != nil {
171 return nil, err
172 }
173
174
175 st.writeVarint(int64(frameTypeHeaders))
176 st.writeVarint(int64(len(headers)))
177 st.Write(headers)
178 if err := st.Flush(); err != nil {
179 return nil, err
180 }
181
182 var bodyAndTrailerWritten bool
183 is100ContinueReq := httpguts.HeaderValuesContainsToken(req.Header["Expect"], "100-continue")
184 if is100ContinueReq {
185 rt.maybeCallWait100Continue()
186 } else {
187 bodyAndTrailerWritten = true
188 wg.Go(func() { cc.writeBodyAndTrailer(rt, req) })
189 }
190
191
192 for {
193 ftype, err := st.readFrameHeader()
194 if err != nil {
195 return nil, err
196 }
197 switch ftype {
198 case frameTypeHeaders:
199 statusCode, h, err := cc.handleHeaders(st)
200 if err != nil {
201 return nil, err
202 }
203
204 if isInfoStatus(statusCode) {
205 if err := rt.maybeCallGot1xxResponse(statusCode, h); err != nil {
206 return nil, err
207 }
208 if statusCode == 100 {
209 rt.maybeCallGot100Continue()
210 if is100ContinueReq && !bodyAndTrailerWritten {
211 bodyAndTrailerWritten = true
212 wg.Go(func() { cc.writeBodyAndTrailer(rt, req) })
213 }
214 }
215 continue
216 }
217
218
219
220 contentLength, err := parseResponseContentLength(req.Method, statusCode, h)
221 if err != nil {
222 return nil, err
223 }
224
225 trailer := make(http.Header)
226 extractTrailerFromHeader(h, trailer)
227 delete(h, "Trailer")
228
229 if (contentLength != 0 && req.Method != http.MethodHead) || len(trailer) > 0 {
230 rt.respBody = &bodyReader{
231 st: st,
232 remain: contentLength,
233 trailer: trailer,
234 }
235 } else {
236 rt.respBody = http.NoBody
237 }
238 resp := &http.Response{
239 Proto: "HTTP/3.0",
240 ProtoMajor: 3,
241 Header: h,
242 StatusCode: statusCode,
243 Status: strconv.Itoa(statusCode) + " " + http.StatusText(statusCode),
244 ContentLength: contentLength,
245 Trailer: trailer,
246 Body: (*transportResponseBody)(rt),
247 }
248 if addedGzip && ascii.EqualFold(h.Get("Content-Encoding"), "gzip") {
249 resp.Body = &httpcommon.GzipReader{Body: resp.Body}
250 h.Del("Content-Encoding")
251 h.Del("Content-Length")
252 resp.ContentLength = -1
253 resp.Uncompressed = true
254 }
255 return resp, nil
256 case frameTypePushPromise:
257 if err := cc.handlePushPromise(st); err != nil {
258 return nil, err
259 }
260 default:
261 if err := st.discardUnknownFrame(ftype); err != nil {
262 return nil, err
263 }
264 }
265 }
266 }
267
268
269
270 func actualContentLength(req *http.Request) int64 {
271 if req.Body == nil || req.Body == http.NoBody {
272 return 0
273 }
274 if req.ContentLength != 0 {
275 return req.ContentLength
276 }
277 return -1
278 }
279
280
281
282
283
284
285
286 func reqBodyIgnored(err error) bool {
287 if streamErr, ok := errors.AsType[quic.StreamError](err); ok {
288 return http3Error(streamErr) == errH3NoError
289 }
290 return false
291 }
292
293
294
295 func (cc *clientConn) writeBodyAndTrailer(rt *roundTripState, req *http.Request) {
296 defer rt.closeReqBody()
297
298 declaredTrailer := req.Trailer.Clone()
299
300 rt.reqBodyWriter.st = rt.st
301 rt.reqBodyWriter.remain = actualContentLength(req)
302 rt.reqBodyWriter.flush = true
303 rt.reqBodyWriter.name = "request"
304 rt.reqBodyWriter.trailer = req.Trailer
305 rt.reqBodyWriter.enc = &cc.enc
306
307 if _, err := io.Copy(&rt.reqBodyWriter, rt.reqBody); err != nil {
308 if reqBodyIgnored(err) {
309 return
310 }
311 rt.abort(err)
312 return
313 }
314
315
316
317 for name := range req.Trailer {
318 if _, ok := declaredTrailer[name]; !ok {
319 delete(req.Trailer, name)
320 }
321 }
322 if err := rt.reqBodyWriter.Close(); err != nil {
323 rt.abort(err)
324 }
325 }
326
327
328 type transportResponseBody roundTripState
329
330
331 func (b *transportResponseBody) Read(p []byte) (n int, err error) {
332 return b.respBody.Read(p)
333 }
334
335 var errRespBodyClosed = errors.New("response body closed")
336
337
338
339 func (b *transportResponseBody) Close() error {
340 rt := (*roundTripState)(b)
341
342 err := rt.respBody.Close()
343 if err == nil {
344 err = errRespBodyClosed
345 }
346 err = rt.abort(err)
347
348
349 rt.closeReqBody()
350 if err == errRespBodyClosed {
351
352
353 return nil
354 }
355 return err
356 }
357
358 func parseResponseContentLength(method string, statusCode int, h http.Header) (int64, error) {
359 clens := h["Content-Length"]
360 if len(clens) == 0 {
361 return -1, nil
362 }
363
364
365
366 for _, v := range clens[1:] {
367 if clens[0] != v {
368 return -1, &streamError{errH3MessageError, "mismatching Content-Length headers"}
369 }
370 }
371
372
373
374
375
376
377 if (statusCode >= 100 && statusCode < 200) ||
378 statusCode == 204 ||
379 (method == "CONNECT" && statusCode >= 200 && statusCode < 300) {
380
381
382 return -1, nil
383 }
384
385 contentLen, err := strconv.ParseUint(clens[0], 10, 63)
386 if err != nil {
387 return -1, &streamError{errH3MessageError, "invalid Content-Length header"}
388 }
389 return int64(contentLen), nil
390 }
391
392 func (cc *clientConn) handleHeaders(st *stream) (statusCode int, h http.Header, err error) {
393 haveStatus := false
394 cookie := ""
395
396
397 err = cc.dec.decode(st, func(_ indexType, name, value string) error {
398 if !httpguts.ValidHeaderFieldValue(value) {
399 return &streamError{errH3MessageError, "invalid field value"}
400 }
401 switch {
402 case name == ":status":
403 if haveStatus {
404 return &streamError{errH3MessageError, "duplicate :status"}
405 }
406 haveStatus = true
407 statusCode, err = strconv.Atoi(value)
408 if err != nil {
409 return &streamError{errH3MessageError, "invalid :status"}
410 }
411 case name[0] == ':':
412
413
414
415
416 return &streamError{errH3MessageError, "undefined pseudo-header"}
417 case name == "cookie":
418
419
420
421
422 if cookie == "" {
423 cookie = value
424 } else {
425 cookie += "; " + value
426 }
427 default:
428 if !validWireHeaderFieldName(name) {
429 return &streamError{errH3MessageError, "invalid field name"}
430 }
431 if h == nil {
432 h = make(http.Header)
433 }
434
435
436
437 cname := httpcommon.CanonicalHeader(name)
438
439
440
441
442
443
444 h[cname] = append(h[cname], value)
445 }
446 return nil
447 })
448 if !haveStatus {
449
450
451 err = errH3MessageError
452 }
453 if cookie != "" {
454 if h == nil {
455 h = make(http.Header)
456 }
457 h["Cookie"] = []string{cookie}
458 }
459 if err := st.endFrame(); err != nil {
460 return 0, nil, err
461 }
462 return statusCode, h, err
463 }
464
465 func (cc *clientConn) handlePushPromise(st *stream) error {
466
467
468
469 return &connectionError{
470 code: errH3IDError,
471 message: "PUSH_PROMISE received when no MAX_PUSH_ID has been sent",
472 }
473 }
474
View as plain text