cloudflare/cloudflared

Public

mirrored from https://github.com/cloudflare/cloudflaredAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
2021.1.4

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

connection/h2mux.go

223lines · modecode

1package connection
2
3import (
4 "context"
5 "net"
6 "net/http"
7 "time"
8
9 "github.com/cloudflare/cloudflared/h2mux"
10 tunnelpogs "github.com/cloudflare/cloudflared/tunnelrpc/pogs"
11 "github.com/cloudflare/cloudflared/websocket"
12
13 "github.com/pkg/errors"
14 "github.com/rs/zerolog"
15 "golang.org/x/sync/errgroup"
16)
17
18const (
19 muxerTimeout = 5 * time.Second
20 openStreamTimeout = 30 * time.Second
21)
22
23type h2muxConnection struct {
24 config *Config
25 muxerConfig *MuxerConfig
26 muxer *h2mux.Muxer
27 // connectionID is only used by metrics, and prometheus requires labels to be string
28 connIndexStr string
29 connIndex uint8
30
31 observer *Observer
32}
33
34type MuxerConfig struct {
35 HeartbeatInterval time.Duration
36 MaxHeartbeats uint64
37 CompressionSetting h2mux.CompressionSetting
38 MetricsUpdateFreq time.Duration
39}
40
41func (mc *MuxerConfig) H2MuxerConfig(h h2mux.MuxedStreamHandler, log *zerolog.Logger) *h2mux.MuxerConfig {
42 return &h2mux.MuxerConfig{
43 Timeout: muxerTimeout,
44 Handler: h,
45 IsClient: true,
46 HeartbeatInterval: mc.HeartbeatInterval,
47 MaxHeartbeats: mc.MaxHeartbeats,
48 Log: log,
49 CompressionQuality: mc.CompressionSetting,
50 }
51}
52
53// NewTunnelHandler returns a TunnelHandler, origin LAN IP and error
54func NewH2muxConnection(
55 config *Config,
56 muxerConfig *MuxerConfig,
57 edgeConn net.Conn,
58 connIndex uint8,
59 observer *Observer,
60) (*h2muxConnection, error, bool) {
61 h := &h2muxConnection{
62 config: config,
63 muxerConfig: muxerConfig,
64 connIndexStr: uint8ToString(connIndex),
65 connIndex: connIndex,
66 observer: observer,
67 }
68
69 // Establish a muxed connection with the edge
70 // Client mux handshake with agent server
71 muxer, err := h2mux.Handshake(edgeConn, edgeConn, *muxerConfig.H2MuxerConfig(h, observer.log), h2mux.ActiveStreams)
72 if err != nil {
73 recoverable := isHandshakeErrRecoverable(err, connIndex, observer)
74 return nil, err, recoverable
75 }
76 h.muxer = muxer
77 return h, nil, false
78}
79
80func (h *h2muxConnection) ServeNamedTunnel(ctx context.Context, namedTunnel *NamedTunnelConfig, credentialManager CredentialManager, connOptions *tunnelpogs.ConnectionOptions, connectedFuse ConnectedFuse) error {
81 errGroup, serveCtx := errgroup.WithContext(ctx)
82 errGroup.Go(func() error {
83 return h.serveMuxer(serveCtx)
84 })
85
86 errGroup.Go(func() error {
87 stream, err := h.newRPCStream(serveCtx, register)
88 if err != nil {
89 return err
90 }
91 rpcClient := newRegistrationRPCClient(ctx, stream, h.observer.log)
92 defer rpcClient.Close()
93
94 if err = rpcClient.RegisterConnection(serveCtx, namedTunnel, connOptions, h.connIndex, h.observer); err != nil {
95 return err
96 }
97 connectedFuse.Connected()
98 return nil
99 })
100
101 errGroup.Go(func() error {
102 h.controlLoop(serveCtx, connectedFuse, true)
103 return nil
104 })
105 return errGroup.Wait()
106}
107
108func (h *h2muxConnection) ServeClassicTunnel(ctx context.Context, classicTunnel *ClassicTunnelConfig, credentialManager CredentialManager, registrationOptions *tunnelpogs.RegistrationOptions, connectedFuse ConnectedFuse) error {
109 errGroup, serveCtx := errgroup.WithContext(ctx)
110 errGroup.Go(func() error {
111 return h.serveMuxer(serveCtx)
112 })
113
114 errGroup.Go(func() (err error) {
115 defer func() {
116 if err == nil {
117 connectedFuse.Connected()
118 }
119 }()
120 if classicTunnel.UseReconnectToken && connectedFuse.IsConnected() {
121 err := h.reconnectTunnel(ctx, credentialManager, classicTunnel, registrationOptions)
122 if err == nil {
123 return nil
124 }
125 // log errors and proceed to RegisterTunnel
126 h.observer.log.Err(err).
127 Uint8(LogFieldConnIndex, h.connIndex).
128 Msg("Couldn't reconnect connection. Re-registering it instead.")
129 }
130 return h.registerTunnel(ctx, credentialManager, classicTunnel, registrationOptions)
131 })
132
133 errGroup.Go(func() error {
134 h.controlLoop(serveCtx, connectedFuse, false)
135 return nil
136 })
137 return errGroup.Wait()
138}
139
140func (h *h2muxConnection) serveMuxer(ctx context.Context) error {
141 // All routines should stop when muxer finish serving. When muxer is shutdown
142 // gracefully, it doesn't return an error, so we need to return errMuxerShutdown
143 // here to notify other routines to stop
144 err := h.muxer.Serve(ctx)
145 if err == nil {
146 return muxerShutdownError{}
147 }
148 return err
149}
150
151func (h *h2muxConnection) controlLoop(ctx context.Context, connectedFuse ConnectedFuse, isNamedTunnel bool) {
152 updateMetricsTickC := time.Tick(h.muxerConfig.MetricsUpdateFreq)
153 for {
154 select {
155 case <-ctx.Done():
156 // UnregisterTunnel blocks until the RPC call returns
157 if connectedFuse.IsConnected() {
158 h.unregister(isNamedTunnel)
159 }
160 h.muxer.Shutdown()
161 return
162 case <-updateMetricsTickC:
163 h.observer.metrics.updateMuxerMetrics(h.connIndexStr, h.muxer.Metrics())
164 }
165 }
166}
167
168func (h *h2muxConnection) newRPCStream(ctx context.Context, rpcName rpcName) (*h2mux.MuxedStream, error) {
169 openStreamCtx, openStreamCancel := context.WithTimeout(ctx, openStreamTimeout)
170 defer openStreamCancel()
171 stream, err := h.muxer.OpenRPCStream(openStreamCtx)
172 if err != nil {
173 return nil, err
174 }
175 return stream, nil
176}
177
178func (h *h2muxConnection) ServeStream(stream *h2mux.MuxedStream) error {
179 respWriter := &h2muxRespWriter{stream}
180
181 req, reqErr := h.newRequest(stream)
182 if reqErr != nil {
183 respWriter.WriteErrorResponse()
184 return reqErr
185 }
186
187 err := h.config.OriginClient.Proxy(respWriter, req, websocket.IsWebSocketUpgrade(req))
188 if err != nil {
189 respWriter.WriteErrorResponse()
190 return err
191 }
192 return nil
193}
194
195func (h *h2muxConnection) newRequest(stream *h2mux.MuxedStream) (*http.Request, error) {
196 req, err := http.NewRequest("GET", "http://localhost:8080", h2mux.MuxedStreamReader{MuxedStream: stream})
197 if err != nil {
198 return nil, errors.Wrap(err, "Unexpected error from http.NewRequest")
199 }
200 err = h2mux.H2RequestHeadersToH1Request(stream.Headers, req)
201 if err != nil {
202 return nil, errors.Wrap(err, "invalid request received")
203 }
204 return req, nil
205}
206
207type h2muxRespWriter struct {
208 *h2mux.MuxedStream
209}
210
211func (rp *h2muxRespWriter) WriteRespHeaders(resp *http.Response) error {
212 headers := h2mux.H1ResponseToH2ResponseHeaders(resp)
213 headers = append(headers, h2mux.Header{Name: ResponseMetaHeaderField, Value: responseMetaHeaderOrigin})
214 return rp.WriteHeaders(headers)
215}
216
217func (rp *h2muxRespWriter) WriteErrorResponse() {
218 _ = rp.WriteHeaders([]h2mux.Header{
219 {Name: ":status", Value: "502"},
220 {Name: ResponseMetaHeaderField, Value: responseMetaHeaderCfd},
221 })
222 _, _ = rp.Write([]byte("502 Bad Gateway"))
223}
224