1
2
3
4
5 package http3
6
7 import (
8 "context"
9 "crypto/tls"
10 "errors"
11 "fmt"
12 "math"
13 "net"
14 "net/http"
15 "net/url"
16 "sync"
17
18 "golang.org/x/net/quic"
19 )
20
21
22
23
24
25
26
27
28 type transport struct {
29 tr1 *http.Transport
30 opts TransportOpts
31
32 mu sync.Mutex
33
34
35
36
37 endpoint *quic.Endpoint
38 activeConns map[*clientConn]struct{}
39 inFlightDials int
40 }
41
42
43
44 type netHTTPTransport struct {
45 *transport
46 }
47
48
49 func (t netHTTPTransport) Registered(tr1 *http.Transport) {
50 t.transport.tr1 = tr1
51 }
52
53
54
55
56
57 func (t netHTTPTransport) RoundTrip(*http.Request) (*http.Response, error) {
58 panic("netHTTPTransport.RoundTrip should never be called")
59 }
60
61 func (t netHTTPTransport) DialClientConn(ctx context.Context, addr string, _ *url.URL, tlsConfig *tls.Config, stateHook func()) (http.RoundTripper, error) {
62 return t.transport.dial(ctx, addr, tlsConfig, stateHook)
63 }
64
65 type TransportOpts struct {
66
67
68
69 ListenQUIC func(addr string, config *quic.Config) (*quic.Endpoint, error)
70
71
72
73
74
75 ListenPacket func(network, addr string) (net.PacketConn, error)
76
77
78
79
80
81
82
83 QUICConfig *quic.Config
84 }
85
86
87 func RegisterTransport(tr *http.Transport, opts TransportOpts) error {
88 tr3 := &transport{
89 opts: opts,
90 activeConns: make(map[*clientConn]struct{}),
91 }
92
93 tr.RegisterProtocol("http/3", netHTTPTransport{tr3})
94 if tr3.tr1 != tr {
95 return errors.New("http3: net/http does not support HTTP/3")
96 }
97 return nil
98 }
99
100 func (tr *transport) incInFlightDials() {
101 tr.mu.Lock()
102 defer tr.mu.Unlock()
103 tr.inFlightDials++
104 }
105
106 func (tr *transport) decInFlightDials() {
107 tr.mu.Lock()
108 defer tr.mu.Unlock()
109 tr.inFlightDials--
110 }
111
112 func (tr *transport) initEndpoint() (err error) {
113 tr.mu.Lock()
114 defer tr.mu.Unlock()
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140 if tr.endpoint == nil {
141 quicConfig := newQUICConfig(tr.opts.QUICConfig, tr.tr1.TLSClientConfig)
142 if tr.opts.ListenQUIC != nil {
143 tr.endpoint, err = tr.opts.ListenQUIC(":0", quicConfig)
144 } else if tr.opts.ListenPacket != nil {
145 var conn net.PacketConn
146 conn, err = tr.opts.ListenPacket("udp", ":0")
147 if err != nil {
148 return err
149 }
150 tr.endpoint, err = quic.NewEndpoint(conn, quicConfig)
151 if err != nil {
152 conn.Close()
153 }
154 } else {
155 tr.endpoint, err = quic.Listen("udp", ":0", quicConfig)
156 }
157 }
158 return err
159 }
160
161
162 func (tr *transport) dial(ctx context.Context, target string, tlsConfig *tls.Config, stateHook func()) (*clientConn, error) {
163 tr.incInFlightDials()
164 defer tr.decInFlightDials()
165
166 if err := tr.initEndpoint(); err != nil {
167 return nil, err
168 }
169 qconn, err := tr.endpoint.Dial(ctx, "udp", target, newQUICConfig(tr.opts.QUICConfig, tlsConfig))
170 if err != nil {
171 return nil, err
172 }
173 return tr.newClientConn(ctx, qconn, stateHook)
174 }
175
176
177
178
179
180
181
182 func (tr *transport) CloseIdleConnections() {
183 tr.mu.Lock()
184 defer tr.mu.Unlock()
185 if tr.endpoint == nil || len(tr.activeConns) > 0 || tr.inFlightDials > 0 {
186 return
187 }
188 tr.endpoint.Close(canceledCtx)
189 tr.endpoint = nil
190 }
191
192
193
194
195 type clientConn struct {
196 tr *transport
197 unregistered chan struct{}
198
199 qconn *quic.Conn
200 genericConn
201
202 enc qpackEncoder
203 dec qpackDecoder
204
205
206 reserved int
207 active int
208 closed bool
209
210 stateHook func()
211 }
212
213 func (tr *transport) registerConn(cc *clientConn) {
214 tr.mu.Lock()
215 defer tr.mu.Unlock()
216 tr.activeConns[cc] = struct{}{}
217 }
218
219 func (tr *transport) unregisterConn(cc *clientConn) {
220 tr.mu.Lock()
221 defer tr.mu.Unlock()
222 delete(tr.activeConns, cc)
223 close(cc.unregistered)
224 }
225
226 func (tr *transport) newClientConn(ctx context.Context, qconn *quic.Conn, stateHook func()) (*clientConn, error) {
227 cc := &clientConn{
228 tr: tr,
229 unregistered: make(chan struct{}),
230 qconn: qconn,
231 stateHook: stateHook,
232 }
233 tr.registerConn(cc)
234 cc.enc.init()
235
236
237 controlStream, err := newConnStream(ctx, cc.qconn, streamTypeControl)
238 if err != nil {
239 tr.unregisterConn(cc)
240 return nil, fmt.Errorf("http3: cannot create control stream: %v", err)
241 }
242 controlStream.writeSettings()
243 controlStream.Flush()
244
245 go func() {
246 cc.acceptStreams(qconn, cc)
247 cc.mu.Lock()
248 cc.closed = true
249 cc.mu.Unlock()
250 cc.maybeCallStateHook()
251 tr.unregisterConn(cc)
252 }()
253 return cc, nil
254 }
255
256 func (cc *clientConn) Close() error {
257 err := cc.qconn.Close()
258
259
260
261
262
263 <-cc.unregistered
264 return err
265 }
266
267 func (cc *clientConn) Err() error {
268 cc.mu.Lock()
269 defer cc.mu.Unlock()
270 if cc.closed {
271 return errors.New("connection closed")
272 }
273 return nil
274 }
275
276 func (cc *clientConn) Reserve() error {
277 cc.mu.Lock()
278 defer cc.mu.Unlock()
279 if cc.closed {
280 return errors.New("connection closed")
281 }
282 cc.reserved++
283 return nil
284 }
285
286 func (cc *clientConn) Release() {
287 cc.mu.Lock()
288 defer cc.mu.Unlock()
289
290
291 if cc.reserved > 0 {
292 cc.reserved--
293 }
294 }
295
296 func (cc *clientConn) Available() int {
297 cc.mu.Lock()
298 defer cc.mu.Unlock()
299 if cc.closed {
300 return 0
301 }
302
303
304
305
306
307
308
309 return math.MaxInt
310 }
311
312 func (cc *clientConn) InFlight() int {
313 cc.mu.Lock()
314 defer cc.mu.Unlock()
315 if cc.closed {
316 return 0
317 }
318 return cc.reserved + cc.active
319 }
320
321 func (cc *clientConn) maybeCallStateHook() {
322 if cc.stateHook != nil {
323 cc.stateHook()
324 }
325 }
326
327 func (cc *clientConn) handleControlStream(st *stream) error {
328
329
330 if err := st.readSettings(func(settingsType, settingsValue int64) error {
331 switch settingsType {
332 case settingsMaxFieldSectionSize:
333 _ = settingsValue
334 case settingsQPACKMaxTableCapacity:
335 _ = settingsValue
336 case settingsQPACKBlockedStreams:
337 _ = settingsValue
338 default:
339
340 }
341 return nil
342 }); err != nil {
343 return err
344 }
345
346 for {
347 ftype, err := st.readFrameHeader()
348 if err != nil {
349 return err
350 }
351 switch ftype {
352 case frameTypeCancelPush:
353
354
355
356
357 return &connectionError{
358 code: errH3IDError,
359 message: "CANCEL_PUSH received when no MAX_PUSH_ID has been sent",
360 }
361 case frameTypeGoaway:
362
363 return errH3NoError
364 default:
365
366 if err := st.discardUnknownFrame(ftype); err != nil {
367 return err
368 }
369 }
370 }
371 }
372
373 func (cc *clientConn) handleEncoderStream(*stream) error {
374
375 return nil
376 }
377
378 func (cc *clientConn) handleDecoderStream(*stream) error {
379
380 return nil
381 }
382
383 func (cc *clientConn) handlePushStream(*stream) error {
384
385
386
387 return &connectionError{
388 code: errH3IDError,
389 message: "push stream created when no MAX_PUSH_ID has been sent",
390 }
391 }
392
393 func (cc *clientConn) handleRequestStream(st *stream) error {
394
395
396
397 return &connectionError{
398 code: errH3StreamCreationError,
399 message: "server created bidirectional stream",
400 }
401 }
402
403
404 func (cc *clientConn) abort(err error) {
405 if e, ok := err.(*connectionError); ok {
406 cc.qconn.Abort(&quic.ConnectionCloseError{
407 Code: uint64(e.code),
408 Reason: e.message,
409 })
410 } else {
411 cc.qconn.Abort(err)
412 }
413 }
414
View as plain text