File
Blob: gateway/proxy_handler.go
| 1 | package gateway |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "errors" |
| 6 | "fmt" |
| 7 | "io" |
| 8 | "log" |
| 9 | "net" |
| 10 | "net/http" |
| 11 | "net/http/httputil" |
| 12 | "slices" |
| 13 | "strings" |
| 14 | "time" |
| 15 | |
| 16 | "go.miragespace.co/specter/spec/protocol" |
| 17 | "go.miragespace.co/specter/spec/rpc" |
| 18 | "go.miragespace.co/specter/spec/tun" |
| 19 | "go.miragespace.co/specter/util" |
| 20 | |
| 21 | "github.com/alecthomas/units" |
| 22 | "github.com/go-chi/chi/v5" |
| 23 | "go.uber.org/zap" |
| 24 | ) |
| 25 | |
| 26 | var delHeaders = []string{ |
| 27 | "True-Client-IP", |
| 28 | "X-Real-IP", |
| 29 | "X-Forwarded-For", |
| 30 | } |
| 31 | |
| 32 | const ( |
| 33 | proxyHeaderTimeout = time.Second * 300 // global limit of how long the gateway will wait for response |
| 34 | ) |
| 35 | |
| 36 | // inspiration from https://blog.cloudflare.com/eliminating-cold-starts-with-cloudflare-workers/ |
| 37 | // warm the route cache when tls handshake begins |
| 38 | func (g *Gateway) HandshakeEarlyHint(sni string) { |
| 39 | if g.HandshakeHintFunc == nil { |
| 40 | return |
| 41 | } |
| 42 | hostname, err := g.extractHostname(sni) |
| 43 | if err != nil { |
| 44 | return |
| 45 | } |
| 46 | if slices.Contains(g.RootDomains, hostname) { |
| 47 | return |
| 48 | } |
| 49 | g.Logger.Debug("Handshake early hint", zap.String("hostname", hostname)) |
| 50 | // TODO: implement a limiter |
| 51 | go g.HandshakeHintFunc(hostname) |
| 52 | } |
| 53 | |
| 54 | func (g *Gateway) httpConnect(w http.ResponseWriter, r *http.Request) { |
| 55 | remote, err := g.connectDialer(r.Context(), r) |
| 56 | if err != nil { |
| 57 | http.Error(w, err.Error(), http.StatusNotFound) |
| 58 | return |
| 59 | } |
| 60 | status := protocol.TunnelStatus{} |
| 61 | err = rpc.BoundedReceive(remote, &status, 1024) |
| 62 | if err != nil { |
| 63 | http.Error(w, err.Error(), http.StatusBadGateway) |
| 64 | remote.Close() |
| 65 | return |
| 66 | } |
| 67 | if status.Status != protocol.TunnelStatusCode_STATUS_OK { |
| 68 | http.Error(w, status.GetError(), http.StatusServiceUnavailable) |
| 69 | remote.Close() |
| 70 | return |
| 71 | } |
| 72 | |
| 73 | rc := http.NewResponseController(w) |
| 74 | local, rw, err := rc.Hijack() |
| 75 | if err != nil { |
| 76 | http.Error(w, err.Error(), http.StatusInternalServerError) |
| 77 | remote.Close() |
| 78 | return |
| 79 | } |
| 80 | |
| 81 | connectResp := http.Response{ |
| 82 | StatusCode: http.StatusOK, |
| 83 | Proto: r.Proto, |
| 84 | ProtoMajor: r.ProtoMajor, |
| 85 | ProtoMinor: r.ProtoMinor, |
| 86 | } |
| 87 | connectResp.Write(rw) |
| 88 | rw.Flush() |
| 89 | |
| 90 | tun.Pipe(local, remote) |
| 91 | } |
| 92 | |
| 93 | func (g *Gateway) extractHostname(host string) (hostname string, err error) { |
| 94 | if net.ParseIP(host) != nil { |
| 95 | err = fmt.Errorf("gateway: hostname cannot be IP") |
| 96 | return |
| 97 | } |
| 98 | if strings.Count(host, ".") < 2 { |
| 99 | err = fmt.Errorf("gateway: too few labels in hostname") |
| 100 | return |
| 101 | } |
| 102 | parts := strings.SplitN(host, ".", 2) |
| 103 | if len(parts) != 2 { |
| 104 | err = fmt.Errorf("gateway: invalid hostname for forwarding") |
| 105 | return |
| 106 | } |
| 107 | if slices.Contains(g.RootDomains, parts[1]) { |
| 108 | hostname = parts[0] |
| 109 | } else { |
| 110 | hostname = host |
| 111 | } |
| 112 | hostname = strings.ToLower(hostname) |
| 113 | return |
| 114 | } |
| 115 | |
| 116 | func (g *Gateway) parseAddr(addr string) (host, hostname string, err error) { |
| 117 | host, _, err = net.SplitHostPort(addr) |
| 118 | if err != nil { |
| 119 | return |
| 120 | } |
| 121 | hostname, err = g.extractHostname(host) |
| 122 | return |
| 123 | } |
| 124 | |
| 125 | func (g *Gateway) connectDialer(ctx context.Context, r *http.Request) (conn net.Conn, err error) { |
| 126 | defer func() { |
| 127 | if pr := recover(); pr != nil { |
| 128 | err = fmt.Errorf("connectDialer panic: %v", pr) |
| 129 | } |
| 130 | }() |
| 131 | |
| 132 | host, hostname, err := g.parseAddr(r.Host) |
| 133 | if err != nil { |
| 134 | return nil, err |
| 135 | } |
| 136 | connectHostname.Add(hostname, 1) |
| 137 | g.Logger.Debug("Dialing to client via overlay (HTTP Connect)", zap.String("hostname", hostname), zap.String("req.URL.Host", host)) |
| 138 | return g.TunnelServer.DialClient(ctx, &protocol.Link{ |
| 139 | Alpn: protocol.Link_TCP, |
| 140 | Hostname: hostname, |
| 141 | Remote: r.RemoteAddr, |
| 142 | }) |
| 143 | } |
| 144 | |
| 145 | func (g *Gateway) overlayDialer(ctx context.Context, addr string) (conn net.Conn, err error) { |
| 146 | defer func() { |
| 147 | if pr := recover(); pr != nil { |
| 148 | err = fmt.Errorf("overlayDialer panic: %v", pr) |
| 149 | } |
| 150 | }() |
| 151 | |
| 152 | host, hostname, err := g.parseAddr(addr) |
| 153 | if err != nil { |
| 154 | return nil, err |
| 155 | } |
| 156 | g.Logger.Debug("Dialing to client via overlay (HTTP)", zap.String("hostname", hostname), zap.String("req.URL.Host", host)) |
| 157 | return g.TunnelServer.DialClient(ctx, &protocol.Link{ |
| 158 | Alpn: protocol.Link_HTTP, |
| 159 | Hostname: hostname, |
| 160 | Remote: g.TunnelServer.Identity().GetAddress(), // include gateway address as placeholder |
| 161 | }) |
| 162 | } |
| 163 | |
| 164 | func (g *Gateway) forwardTCP(ctx context.Context, host string, remote string, conn DeadlineReadWriteCloser) error { |
| 165 | var ( |
| 166 | hostname string |
| 167 | c net.Conn |
| 168 | err error |
| 169 | ) |
| 170 | defer func() { |
| 171 | if err != nil { |
| 172 | tun.SendStatusProto(conn, err) |
| 173 | conn.Close() |
| 174 | return |
| 175 | } |
| 176 | tun.Pipe(conn, c) |
| 177 | }() |
| 178 | |
| 179 | // because of quic's early connection, the client need to "poke" us before |
| 180 | // we can actually accept a stream, despite .OpenStreamSync |
| 181 | conn.SetReadDeadline(time.Now().Add(time.Second * 3)) |
| 182 | err = tun.DrainStatusProto(conn) |
| 183 | if err != nil { |
| 184 | return err |
| 185 | } |
| 186 | conn.SetReadDeadline(time.Time{}) |
| 187 | |
| 188 | hostname, err = g.extractHostname(host) |
| 189 | if err != nil { |
| 190 | return err |
| 191 | } |
| 192 | |
| 193 | g.Logger.Debug("Dialing to client via overlay (TCP)", zap.String("hostname", host)) |
| 194 | c, err = g.TunnelServer.DialClient(ctx, &protocol.Link{ |
| 195 | Alpn: protocol.Link_TCP, |
| 196 | Hostname: hostname, |
| 197 | Remote: remote, |
| 198 | }) |
| 199 | if err != nil { |
| 200 | return err |
| 201 | } |
| 202 | |
| 203 | return nil |
| 204 | } |
| 205 | |
| 206 | func (g *Gateway) proxyRewrite(preq *httputil.ProxyRequest) { |
| 207 | in := preq.In |
| 208 | out := preq.Out |
| 209 | |
| 210 | out.URL.Scheme = "https" |
| 211 | // Most browsers' connection coalescing behavior for http3 is the same as http2, |
| 212 | // as they could be reusing the same tcp/quic connection for different hosts. |
| 213 | // https://daniel.haxx.se/blog/2016/08/18/http2-connection-coalescing/ |
| 214 | // https://mailarchive.ietf.org/arch/msg/quic/ffjARd8-IobIE2T9_r5u9hBDbuk/ |
| 215 | if in.ProtoAtLeast(2, 0) { |
| 216 | out.URL.Host = in.Host |
| 217 | } else if in.TLS != nil { |
| 218 | out.URL.Host = in.TLS.ServerName |
| 219 | } else { |
| 220 | out.URL.Host = in.Host |
| 221 | } |
| 222 | out.URL.Host = out.URL.Hostname() |
| 223 | out.Host = out.URL.Host |
| 224 | |
| 225 | for _, header := range delHeaders { |
| 226 | out.Header.Del(header) |
| 227 | } |
| 228 | |
| 229 | // This is the public trust boundary. Rewrite has already removed Forwarded |
| 230 | // and the three X-Forwarded fields set below; never restore visitor-supplied |
| 231 | // values. The tunnel carries one visitor IP, without appending tunnel peers. |
| 232 | preq.SetXForwarded() |
| 233 | if g.GatewayPort == 443 { |
| 234 | out.Header.Set("X-Forwarded-Host", out.URL.Host) |
| 235 | } else { |
| 236 | out.Header.Set("X-Forwarded-Host", fmt.Sprintf("%s:%d", out.URL.Host, g.GatewayPort)) |
| 237 | } |
| 238 | out.Header.Set("X-Forwarded-Proto", "https") |
| 239 | } |
| 240 | |
| 241 | func (g *Gateway) proxyHandler(proxyLogger *log.Logger) http.Handler { |
| 242 | respHandler := func(r *http.Response) error { |
| 243 | r.Header.Del("alt-svc") |
| 244 | g.appendHeaders(r.Request.ProtoAtLeast(3, 0))(r.Header) |
| 245 | return nil |
| 246 | } |
| 247 | |
| 248 | g.Logger.Info("Using buffer sizes for proxy", |
| 249 | zap.String("transport", units.Base2Bytes(g.Options.TransportBufferSize).String()), |
| 250 | zap.String("proxy", units.Base2Bytes(g.Options.ProxyBufferSize).String()), |
| 251 | ) |
| 252 | |
| 253 | proxyTransport := &http.Transport{ |
| 254 | MaxConnsPerHost: 30, |
| 255 | MaxIdleConnsPerHost: 5, |
| 256 | DisableCompression: true, |
| 257 | IdleConnTimeout: time.Second * 30, |
| 258 | ResponseHeaderTimeout: proxyHeaderTimeout, |
| 259 | ExpectContinueTimeout: time.Second * 5, |
| 260 | WriteBufferSize: g.Options.TransportBufferSize, |
| 261 | ReadBufferSize: g.Options.TransportBufferSize, |
| 262 | } |
| 263 | proxyTransport.DialTLSContext = func(ctx context.Context, _, addr string) (net.Conn, error) { |
| 264 | return g.overlayDialer(ctx, addr) |
| 265 | } |
| 266 | |
| 267 | bufPool := util.NewBufferPool(g.Options.ProxyBufferSize) |
| 268 | proxy := &httputil.ReverseProxy{ |
| 269 | Rewrite: g.proxyRewrite, |
| 270 | Transport: proxyTransport, |
| 271 | BufferPool: bufPool, |
| 272 | ErrorHandler: g.errorHandler, |
| 273 | ModifyResponse: respHandler, |
| 274 | ErrorLog: proxyLogger, |
| 275 | FlushInterval: -1, |
| 276 | } |
| 277 | |
| 278 | router := chi.NewRouter() |
| 279 | |
| 280 | g.mountCgiHandler(router) |
| 281 | router.Handle("/*", proxy) |
| 282 | |
| 283 | return router |
| 284 | } |
| 285 | |
| 286 | func (g *Gateway) errorHandler(w http.ResponseWriter, r *http.Request, e error) { |
| 287 | g.appendHeaders(r.ProtoAtLeast(3, 0))(w.Header()) |
| 288 | |
| 289 | if errors.Is(e, tun.ErrDestinationNotFound) { |
| 290 | w.WriteHeader(http.StatusNotFound) |
| 291 | fmt.Fprintf(w, "Destination %s not found on the specter network.", r.URL.Hostname()) |
| 292 | return |
| 293 | } |
| 294 | |
| 295 | if errors.Is(e, tun.ErrTunnelClientNotConnected) { |
| 296 | w.WriteHeader(http.StatusServiceUnavailable) |
| 297 | fmt.Fprintf(w, "Destination %s is not connected to specter network.", r.URL.Hostname()) |
| 298 | return |
| 299 | } |
| 300 | |
| 301 | if errors.Is(e, context.Canceled) || |
| 302 | errors.Is(e, io.EOF) { |
| 303 | // this is expected |
| 304 | return |
| 305 | } |
| 306 | |
| 307 | logger := g.Logger.With(zap.String("hostname", r.URL.Hostname())) |
| 308 | |
| 309 | logger.Debug("error forwarding http request", zap.Error(e)) |
| 310 | |
| 311 | if tun.IsTimeout(e) { |
| 312 | w.WriteHeader(http.StatusGatewayTimeout) |
| 313 | fmt.Fprintf(w, "Destination %s is taking too long to respond.", r.URL.Hostname()) |
| 314 | return |
| 315 | } |
| 316 | |
| 317 | logger.Error("error forwarding to client", zap.Error(e)) |
| 318 | w.WriteHeader(http.StatusBadGateway) |
| 319 | fmt.Fprint(w, "An unexpected error has occurred while attempting to forward to destination.") |
| 320 | } |