Skip to content
File

Blob: tun/client/proxy.go

go231 lines
1package client
2 
3import (
4 "context"
5 "crypto/tls"
6 "errors"
7 "fmt"
8 "io"
9 "net"
10 "net/http"
11 "net/http/httputil"
12 "net/url"
13 "strings"
14 "time"
15 
16 "go.miragespace.co/specter/spec/tun"
17 "go.miragespace.co/specter/util"
18 "go.miragespace.co/specter/util/acceptor"
19 "go.miragespace.co/specter/util/pipe"
20 
21 "go.uber.org/zap"
22 "go.uber.org/zap/zapcore"
23)
24 
25type httpProxy struct {
26 acceptor *acceptor.HTTP2Acceptor
27 forwarder *http.Server
28}
29 
30type httpReqCtxKey string
31 
32const (
33 ctxStartTime = httpReqCtxKey("start-time")
34 defaultProxyHeaderTimeout = time.Second * 15
35)
36 
37func injectStartTime(proxy http.Handler) http.Handler {
38 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
39 start := time.Now()
40 ctx := context.WithValue(r.Context(), ctxStartTime, start)
41 r = r.WithContext(ctx)
42 proxy.ServeHTTP(w, r)
43 })
44}
45 
46func (f *forwarder) forwardStream(ctx context.Context, hostname string, remote net.Conn, r route) {
47 var (
48 u *url.URL = r.parsed
49 target string
50 local net.Conn
51 err error
52 )
53 logger := f.logger.With(zap.String("hostname", hostname), zap.String("target", u.String()))
54 switch u.Scheme {
55 case "tcp":
56 dialer := &net.Dialer{
57 Timeout: time.Second * 3,
58 }
59 target = u.Host
60 local, err = dialer.DialContext(ctx, "tcp", u.Host)
61 case "unix", "winio":
62 target = u.Path
63 local, err = pipe.DialPipe(ctx, u.Path)
64 default:
65 err = fmt.Errorf("unknown scheme: %s", u.Scheme)
66 }
67 if err != nil {
68 logger.Error("Error dialing to target", zap.String("target", target), zap.Error(err))
69 tun.SendStatusProto(remote, err)
70 remote.Close()
71 return
72 }
73 tun.SendStatusProto(remote, nil)
74 tun.Pipe(remote, local)
75}
76 
77func (f *forwarder) getHTTPProxy(_ context.Context, hostname string, r route) *httpProxy {
78 proxy, loaded := f.proxies.LoadOrStoreLazy(hostname, func() *httpProxy {
79 var (
80 u *url.URL = r.parsed
81 isPipe = false
82 )
83 
84 logger := f.logger.With(zap.String("hostname", hostname), zap.String("target", u.String()))
85 logger.Info("Creating new proxy")
86 
87 readHeaderTimeout := defaultProxyHeaderTimeout
88 if r.proxyHeaderReadTimeout > 0 {
89 readHeaderTimeout = r.proxyHeaderReadTimeout
90 }
91 
92 tp := &http.Transport{
93 TLSClientConfig: &tls.Config{
94 ServerName: u.Host,
95 InsecureSkipVerify: r.insecure,
96 },
97 MaxIdleConns: 10,
98 IdleConnTimeout: time.Second * 30,
99 ResponseHeaderTimeout: readHeaderTimeout,
100 ForceAttemptHTTP2: true,
101 }
102 switch u.Scheme {
103 case "unix", "winio":
104 isPipe = true
105 tp.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
106 return pipe.DialPipe(ctx, u.Path)
107 }
108 }
109 
110 // determine Host header behavior based on ProxyHeaderMode
111 mode := strings.ToLower(r.proxyHeaderMode)
112 root := f.rootDomain.Load()
113 fqdn := hostname
114 if !strings.Contains(fqdn, ".") && root != "" {
115 fqdn = fmt.Sprintf("%s.%s", fqdn, root)
116 }
117 
118 hostHeader := u.Host
119 if isPipe {
120 hostHeader = "pipe"
121 }
122 
123 switch mode {
124 case "":
125 // legacy behavior: optional override only
126 if len(r.proxyHeaderHost) > 0 {
127 hostHeader = r.proxyHeaderHost
128 }
129 case "hostname":
130 hostHeader = fqdn
131 case "target":
132 if !isPipe {
133 hostHeader = u.Host
134 } else {
135 // validate() should have prevented this; fall back to hostname for safety
136 hostHeader = fqdn
137 }
138 case "custom":
139 if len(r.proxyHeaderHost) > 0 {
140 hostHeader = r.proxyHeaderHost
141 }
142 }
143 
144 proxy := &httputil.ReverseProxy{}
145 proxy.Rewrite = func(preq *httputil.ProxyRequest) {
146 if isPipe {
147 preq.Out.URL.Scheme = "http"
148 preq.Out.URL.Host = "pipe"
149 preq.Out.URL.Path = preq.In.URL.Path
150 preq.Out.URL.RawPath = preq.In.URL.RawPath
151 preq.Out.URL.RawQuery = preq.In.URL.RawQuery
152 } else {
153 preq.SetURL(u)
154 }
155 preq.Out.Host = hostHeader
156 // This private proxy receives only authenticated server streams or
157 // the in-process gateway. Restore the visitor IP, host, and scheme
158 // sanitized by that gateway, without adding the tunnel peer.
159 for _, header := range []string{"X-Forwarded-For", "X-Forwarded-Host", "X-Forwarded-Proto"} {
160 if value := preq.In.Header.Get(header); value != "" {
161 preq.Out.Header.Set(header, value)
162 }
163 }
164 }
165 proxy.Transport = tp
166 proxy.ErrorHandler = func(rw http.ResponseWriter, r *http.Request, e error) {
167 if errors.Is(e, context.Canceled) ||
168 errors.Is(e, io.EOF) {
169 // this is expected
170 return
171 }
172 logger.Error("Error forwarding http/https request", zap.Object("request", (*encRequest)(r)), zap.Error(e))
173 rw.WriteHeader(http.StatusBadGateway)
174 fmt.Fprintf(rw, "Forwarding target returned error: %s", e.Error())
175 }
176 proxy.ModifyResponse = func(r *http.Response) error {
177 logger.Debug("Access Log", zap.Object("request", (*encRequest)(r.Request)), zap.Object("response", (*encResponse)(r)))
178 return nil
179 }
180 proxy.ErrorLog = util.GetStdLogger(logger, "targetProxy")
181 proxy.BufferPool = util.NewBufferPool(1024 * 16)
182 
183 return &httpProxy{
184 acceptor: acceptor.NewH2Acceptor(nil),
185 forwarder: &http.Server{
186 Handler: injectStartTime(proxy),
187 ErrorLog: zap.NewStdLog(f.logger),
188 ReadHeaderTimeout: readHeaderTimeout,
189 },
190 }
191 })
192 if !loaded {
193 go proxy.forwarder.Serve(proxy.acceptor)
194 }
195 return proxy
196}
197 
198type encRequest http.Request
199 
200var _ zapcore.ObjectMarshaler = (*encRequest)(nil)
201 
202func (r *encRequest) MarshalLogObject(enc zapcore.ObjectEncoder) error {
203 enc.AddString("method", r.Method)
204 enc.AddString("path", r.RequestURI)
205 proxied := r.Header.Get("x-forwarded-for")
206 if len(proxied) > 0 {
207 ips := strings.Split(proxied, ",")
208 for i, ip := range ips {
209 ips[i] = strings.TrimSpace(ip)
210 }
211 enc.AddString("client", ips[0])
212 if len(ips) > 1 {
213 enc.AddString("via", ips[1])
214 }
215 }
216 return nil
217}
218 
219type encResponse http.Response
220 
221var _ zapcore.ObjectMarshaler = (*encResponse)(nil)
222 
223func (r *encResponse) MarshalLogObject(enc zapcore.ObjectEncoder) error {
224 ctx := r.Request.Context()
225 start := ctx.Value(ctxStartTime).(time.Time)
226 enc.AddInt("code", r.StatusCode)
227 enc.AddInt64("bytes", r.ContentLength)
228 enc.AddDuration("duration", time.Since(start))
229 return nil
230}