Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 48 additions & 5 deletions internal/http3/transport.go
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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"
Expand Down
29 changes: 27 additions & 2 deletions transport.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down Expand Up @@ -146,6 +147,8 @@ type Transport struct {
httpRoundTripWrappers []HttpRoundTripWrapper
}

type HTTP3DialError = http3.DialError

// NewTransport is an alias of T
func NewTransport() *Transport {
return T()
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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.
Expand Down