diff --git a/internal/http3/transport.go b/internal/http3/transport.go index ddb7c6ed..70bdf06f 100644 --- a/internal/http3/transport.go +++ b/internal/http3/transport.go @@ -76,8 +76,8 @@ type Transport struct { // Dial specifies an optional dial function for creating QUIC // connections for requests. - // If Dial is nil, a UDPConn will be created at the first request - // and will be reused for subsequent connections to other servers. + // If Dial is nil, DialContext is used when set. Otherwise, a UDPConn + // is created at the first request and reused for subsequent connections. Dial func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) // Enable support for HTTP/3 datagrams (RFC 9297). @@ -126,6 +126,30 @@ var ( ErrTransportClosed = errors.New("http3: transport is closed") ) +type DialError struct { + error +} + +func (e *DialError) Unwrap() error { + return e.error +} + +type connectedPacketConn struct { + net.Conn +} + +func (c connectedPacketConn) ReadFrom(buf []byte) (int, net.Addr, error) { + n, err := c.Read(buf) + return n, c.RemoteAddr(), err +} + +func (c connectedPacketConn) WriteTo(buf []byte, addr net.Addr) (int, error) { + if addr.String() != c.RemoteAddr().String() { + return 0, net.InvalidAddrError("connected UDP peer changed") + } + return c.Write(buf) +} + func (t *Transport) init() error { if t.newClientConn == nil { t.newClientConn = func(conn *quic.Conn) clientConn { @@ -159,7 +183,7 @@ func (t *Transport) init() error { if t.QUICConfig.MaxIncomingStreams == 0 { t.QUICConfig.MaxIncomingStreams = -1 // don't allow any bidirectional streams } - if t.Dial == nil { + if t.Dial == nil && (t.Options == nil || t.DialContext == nil) { udpConn, err := net.ListenUDP("udp", nil) if err != nil { return err @@ -366,11 +390,15 @@ func (t *Transport) getClient(ctx context.Context, hostname string, onlyCached b } func (t *Transport) dial(ctx context.Context, hostname string) (*quic.Conn, clientConn, error) { + tlsConfig := t.TLSClientConfig + if tlsConfig == nil && t.Options != nil { + tlsConfig = t.Options.TLSClientConfig + } var tlsConf *tls.Config - if t.TLSClientConfig == nil { + if tlsConfig == nil { tlsConf = &tls.Config{} } else { - tlsConf = t.TLSClientConfig.Clone() + tlsConf = tlsConfig.Clone() } if tlsConf.ServerName == "" { sni, _, err := net.SplitHostPort(hostname) @@ -384,6 +412,21 @@ func (t *Transport) dial(ctx context.Context, hostname string) (*quic.Conn, clie tlsConf.NextProtos = []string{NextProtoH3} dial := t.Dial + if dial == nil && t.Options != nil && t.DialContext != nil { + dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) { + udp, err := t.DialContext(ctx, "udp", addr) + if err != nil { + return nil, &DialError{err} + } + conn, err := quic.Dial(ctx, connectedPacketConn{udp}, udp.RemoteAddr(), tlsCfg, cfg) + if err != nil { + _ = udp.Close() + return nil, &DialError{err} + } + context.AfterFunc(conn.Context(), func() { _ = udp.Close() }) + return conn, nil + } + } if dial == nil { dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) { network := "udp" diff --git a/transport.go b/transport.go index 75ef99f5..c066b397 100644 --- a/transport.go +++ b/transport.go @@ -47,6 +47,7 @@ import ( "github.com/imroc/req/v3/internal/util" "github.com/imroc/req/v3/pkg/altsvc" reqtls "github.com/imroc/req/v3/pkg/tls" + "github.com/quic-go/quic-go" htmlcharset "golang.org/x/net/html/charset" "golang.org/x/text/encoding/ianaindex" @@ -146,6 +147,8 @@ type Transport struct { httpRoundTripWrappers []HttpRoundTripWrapper } +type HTTP3DialError = http3.DialError + // NewTransport is an alias of T func NewTransport() *Transport { return T() @@ -473,8 +476,8 @@ func (t *Transport) SetProxy(proxy func(*http.Request) (*url.URL, error)) *Trans return t } -// SetDial set the custom DialContext function, only valid for HTTP1 and HTTP2, which specifies the -// dial function for creating unencrypted TCP connections. +// SetDial set the custom DialContext function, which specifies the +// dial function for creating TCP or UDP connections. // If it is nil, then the transport dials using package net. // // The dial function runs concurrently with calls to RoundTrip. @@ -555,6 +558,14 @@ func (t *Transport) EnableForceHTTP3() *Transport { return t } +func (t *Transport) SetHTTP3QUICConfig(config *quic.Config) *Transport { + t.EnableHTTP3() + if t.t3 != nil { + t.t3.QUICConfig = config + } + return t +} + // DisableForceHttpVersion disable force using specified http // version (disabled by default). func (t *Transport) DisableForceHttpVersion() *Transport { @@ -776,6 +787,9 @@ func (t *Transport) Clone() *Transport { } if t.t3 != nil { tt.EnableHTTP3() + if t.t3.QUICConfig != nil { + tt.t3.QUICConfig = t.t3.QUICConfig.Clone() + } } return tt } @@ -1225,6 +1239,17 @@ func (t *Transport) CloseIdleConnections() { if t2 := t.t2; t2 != nil { t2.CloseIdleConnections() } + if t3 := t.t3; t3 != nil { + t3.CloseIdleConnections() + } +} + +func (t *Transport) Close() error { + t.CloseIdleConnections() + if t.t3 != nil { + return t.t3.Close() + } + return nil } // prepareTransportCancel sets up state to convert Transport.CancelRequest into context cancellation.