1
2
3
4
5
6
7 package httputil
8
9 import (
10 "context"
11 "errors"
12 "fmt"
13 "internal/godebug"
14 "io"
15 "log"
16 "mime"
17 "net"
18 "net/http"
19 "net/http/httptrace"
20 "net/http/internal/ascii"
21 "net/textproto"
22 "net/url"
23 "strings"
24 "sync"
25 "time"
26
27 "golang.org/x/net/http/httpguts"
28 )
29
30
31 type ProxyRequest struct {
32
33
34 In *http.Request
35
36
37
38
39
40 Out *http.Request
41 }
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57 func (r *ProxyRequest) SetURL(target *url.URL) {
58 rewriteRequestURL(r.Out, target)
59 r.Out.Host = ""
60 }
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81 func (r *ProxyRequest) SetXForwarded() {
82 clientIP, _, err := net.SplitHostPort(r.In.RemoteAddr)
83 if err == nil {
84 prior := r.Out.Header["X-Forwarded-For"]
85 if len(prior) > 0 {
86 clientIP = strings.Join(prior, ", ") + ", " + clientIP
87 }
88 r.Out.Header.Set("X-Forwarded-For", clientIP)
89 } else {
90 r.Out.Header.Del("X-Forwarded-For")
91 }
92 r.Out.Header.Set("X-Forwarded-Host", r.In.Host)
93 if r.In.TLS == nil {
94 r.Out.Header.Set("X-Forwarded-Proto", "http")
95 } else {
96 r.Out.Header.Set("X-Forwarded-Proto", "https")
97 }
98 }
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113 type ReverseProxy struct {
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135 Rewrite func(*ProxyRequest)
136
137
138
139 Transport http.RoundTripper
140
141
142
143
144
145
146
147
148
149
150
151 FlushInterval time.Duration
152
153
154
155
156 ErrorLog *log.Logger
157
158
159
160
161 BufferPool BufferPool
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176 ModifyResponse func(*http.Response) error
177
178
179
180
181
182
183 ErrorHandler func(http.ResponseWriter, *http.Request, error)
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265 Director func(*http.Request)
266 }
267
268
269
270 type BufferPool interface {
271 Get() []byte
272 Put([]byte)
273 }
274
275 func singleJoiningSlash(a, b string) string {
276 aslash := strings.HasSuffix(a, "/")
277 bslash := strings.HasPrefix(b, "/")
278 switch {
279 case aslash && bslash:
280 return a + b[1:]
281 case !aslash && !bslash:
282 return a + "/" + b
283 }
284 return a + b
285 }
286
287 func joinURLPath(a, b *url.URL) (path, rawpath string) {
288 if a.RawPath == "" && b.RawPath == "" {
289 return singleJoiningSlash(a.Path, b.Path), ""
290 }
291
292
293 apath := a.EscapedPath()
294 bpath := b.EscapedPath()
295
296 aslash := strings.HasSuffix(apath, "/")
297 bslash := strings.HasPrefix(bpath, "/")
298
299 switch {
300 case aslash && bslash:
301 return a.Path + b.Path[1:], apath + bpath[1:]
302 case !aslash && !bslash:
303 return a.Path + "/" + b.Path, apath + "/" + bpath
304 }
305 return a.Path + b.Path, apath + bpath
306 }
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332 func NewSingleHostReverseProxy(target *url.URL) *ReverseProxy {
333 director := func(req *http.Request) {
334 rewriteRequestURL(req, target)
335 }
336 return &ReverseProxy{Director: director}
337 }
338
339 func rewriteRequestURL(req *http.Request, target *url.URL) {
340 targetQuery := target.RawQuery
341 req.URL.Scheme = target.Scheme
342 req.URL.Host = target.Host
343 req.URL.Path, req.URL.RawPath = joinURLPath(target, req.URL)
344 if targetQuery == "" || req.URL.RawQuery == "" {
345 req.URL.RawQuery = targetQuery + req.URL.RawQuery
346 } else {
347 req.URL.RawQuery = targetQuery + "&" + req.URL.RawQuery
348 }
349 }
350
351 func copyHeader(dst, src http.Header) {
352 for k, vv := range src {
353 for _, v := range vv {
354 dst.Add(k, v)
355 }
356 }
357 }
358
359
360
361
362
363
364 var hopHeaders = []string{
365 "Connection",
366 "Proxy-Connection",
367 "Keep-Alive",
368 "Proxy-Authenticate",
369 "Proxy-Authorization",
370 "Te",
371 "Trailer",
372 "Transfer-Encoding",
373 "Upgrade",
374 }
375
376 func (p *ReverseProxy) defaultErrorHandler(rw http.ResponseWriter, req *http.Request, err error) {
377 p.logf("http: proxy error: %v", err)
378 rw.WriteHeader(http.StatusBadGateway)
379 }
380
381 func (p *ReverseProxy) getErrorHandler() func(http.ResponseWriter, *http.Request, error) {
382 if p.ErrorHandler != nil {
383 return p.ErrorHandler
384 }
385 return p.defaultErrorHandler
386 }
387
388
389
390 func (p *ReverseProxy) modifyResponse(rw http.ResponseWriter, res *http.Response, req *http.Request) bool {
391 if p.ModifyResponse == nil {
392 return true
393 }
394 if err := p.ModifyResponse(res); err != nil {
395 res.Body.Close()
396 p.getErrorHandler()(rw, req, err)
397 return false
398 }
399 return true
400 }
401
402 func (p *ReverseProxy) ServeHTTP(rw http.ResponseWriter, req *http.Request) {
403 transport := p.Transport
404 if transport == nil {
405 transport = http.DefaultTransport
406 }
407
408 ctx := req.Context()
409 if ctx.Done() != nil {
410
411
412
413
414
415
416
417
418
419
420 } else if cn, ok := rw.(http.CloseNotifier); ok {
421 var cancel context.CancelFunc
422 ctx, cancel = context.WithCancel(ctx)
423 defer cancel()
424 notifyChan := cn.CloseNotify()
425 go func() {
426 select {
427 case <-notifyChan:
428 cancel()
429 case <-ctx.Done():
430 }
431 }()
432 }
433
434 outreq := req.Clone(ctx)
435 if req.ContentLength == 0 {
436 outreq.Body = nil
437 }
438 if outreq.Body != nil {
439
440
441
442
443
444
445 defer outreq.Body.Close()
446 }
447 if outreq.Header == nil {
448 outreq.Header = make(http.Header)
449 }
450
451 if (p.Director != nil) == (p.Rewrite != nil) {
452 p.getErrorHandler()(rw, req, errors.New("ReverseProxy must have exactly one of Director or Rewrite set"))
453 return
454 }
455
456 if p.Director != nil {
457 p.Director(outreq)
458 if outreq.Form != nil {
459 outreq.URL.RawQuery = cleanQueryParams(outreq.URL.RawQuery)
460 }
461 }
462 outreq.Close = false
463
464 reqUpType := upgradeType(outreq.Header)
465 if !ascii.IsPrint(reqUpType) {
466 p.getErrorHandler()(rw, req, fmt.Errorf("client tried to switch to invalid protocol %q", reqUpType))
467 return
468 }
469 removeHopByHopHeaders(outreq.Header)
470
471
472
473
474
475
476 if httpguts.HeaderValuesContainsToken(req.Header["Te"], "trailers") {
477 outreq.Header.Set("Te", "trailers")
478 }
479
480
481
482 if reqUpType != "" {
483 outreq.Header.Set("Connection", "Upgrade")
484 outreq.Header.Set("Upgrade", reqUpType)
485 }
486
487 if p.Rewrite != nil {
488
489
490
491 outreq.Header.Del("Forwarded")
492 outreq.Header.Del("X-Forwarded-For")
493 outreq.Header.Del("X-Forwarded-Host")
494 outreq.Header.Del("X-Forwarded-Proto")
495
496
497 outreq.URL.RawQuery = cleanQueryParams(outreq.URL.RawQuery)
498
499 pr := &ProxyRequest{
500 In: req,
501 Out: outreq,
502 }
503 p.Rewrite(pr)
504 outreq = pr.Out
505 } else {
506 if clientIP, _, err := net.SplitHostPort(req.RemoteAddr); err == nil {
507
508
509
510 prior, ok := outreq.Header["X-Forwarded-For"]
511 omit := ok && prior == nil
512 if len(prior) > 0 {
513 clientIP = strings.Join(prior, ", ") + ", " + clientIP
514 }
515 if !omit {
516 outreq.Header.Set("X-Forwarded-For", clientIP)
517 }
518 }
519 }
520
521 if _, ok := outreq.Header["User-Agent"]; !ok {
522
523
524 outreq.Header.Set("User-Agent", "")
525 }
526
527 var (
528 roundTripMutex sync.Mutex
529 roundTripDone bool
530 )
531 trace := &httptrace.ClientTrace{
532 Got1xxResponse: func(code int, header textproto.MIMEHeader) error {
533 roundTripMutex.Lock()
534 defer roundTripMutex.Unlock()
535 if roundTripDone {
536
537
538 return nil
539 }
540 h := rw.Header()
541 copyHeader(h, http.Header(header))
542 rw.WriteHeader(code)
543
544
545 clear(h)
546 return nil
547 },
548 }
549 outreq = outreq.WithContext(httptrace.WithClientTrace(outreq.Context(), trace))
550
551 res, err := transport.RoundTrip(outreq)
552 roundTripMutex.Lock()
553 roundTripDone = true
554 roundTripMutex.Unlock()
555 if err != nil {
556 p.getErrorHandler()(rw, outreq, err)
557 return
558 }
559
560
561 if res.StatusCode == http.StatusSwitchingProtocols {
562 if !p.modifyResponse(rw, res, outreq) {
563 return
564 }
565 p.handleUpgradeResponse(rw, outreq, res)
566 return
567 }
568
569 removeHopByHopHeaders(res.Header)
570
571 if !p.modifyResponse(rw, res, outreq) {
572 return
573 }
574
575 copyHeader(rw.Header(), res.Header)
576
577
578
579 announcedTrailers := len(res.Trailer)
580 if announcedTrailers > 0 {
581 trailerKeys := make([]string, 0, len(res.Trailer))
582 for k := range res.Trailer {
583 trailerKeys = append(trailerKeys, k)
584 }
585 rw.Header().Add("Trailer", strings.Join(trailerKeys, ", "))
586 }
587
588 rw.WriteHeader(res.StatusCode)
589
590 err = p.copyResponse(rw, res.Body, p.flushInterval(res))
591 if err != nil {
592 defer res.Body.Close()
593
594
595
596 if !shouldPanicOnCopyError(req) {
597 p.logf("suppressing panic for copyResponse error in test; copy error: %v", err)
598 return
599 }
600 panic(http.ErrAbortHandler)
601 }
602 res.Body.Close()
603
604 if len(res.Trailer) > 0 {
605
606
607
608 http.NewResponseController(rw).Flush()
609 }
610
611 if len(res.Trailer) == announcedTrailers {
612 copyHeader(rw.Header(), res.Trailer)
613 return
614 }
615
616 for k, vv := range res.Trailer {
617 k = http.TrailerPrefix + k
618 for _, v := range vv {
619 rw.Header().Add(k, v)
620 }
621 }
622 }
623
624 var inOurTests bool
625
626
627
628
629
630
631 func shouldPanicOnCopyError(req *http.Request) bool {
632 if inOurTests {
633
634 return true
635 }
636 if req.Context().Value(http.ServerContextKey) != nil {
637
638
639 return true
640 }
641
642
643 return false
644 }
645
646
647 func removeHopByHopHeaders(h http.Header) {
648
649 for _, f := range h["Connection"] {
650 for sf := range strings.SplitSeq(f, ",") {
651 if sf = textproto.TrimString(sf); sf != "" {
652 h.Del(sf)
653 }
654 }
655 }
656
657
658
659 for _, f := range hopHeaders {
660 h.Del(f)
661 }
662 }
663
664
665
666 func (p *ReverseProxy) flushInterval(res *http.Response) time.Duration {
667 resCT := res.Header.Get("Content-Type")
668
669
670
671 if baseCT, _, _ := mime.ParseMediaType(resCT); baseCT == "text/event-stream" {
672 return -1
673 }
674
675
676 if res.ContentLength == -1 {
677 return -1
678 }
679
680 return p.FlushInterval
681 }
682
683 func (p *ReverseProxy) copyResponse(dst http.ResponseWriter, src io.Reader, flushInterval time.Duration) error {
684 var w io.Writer = dst
685
686 if flushInterval != 0 {
687 mlw := &maxLatencyWriter{
688 dst: dst,
689 flush: http.NewResponseController(dst).Flush,
690 latency: flushInterval,
691 }
692 defer mlw.stop()
693
694
695 mlw.flushPending = true
696 mlw.t = time.AfterFunc(flushInterval, mlw.delayedFlush)
697
698 w = mlw
699 }
700
701 var buf []byte
702 if p.BufferPool != nil {
703 buf = p.BufferPool.Get()
704 defer p.BufferPool.Put(buf)
705 }
706 _, err := p.copyBuffer(w, src, buf)
707 return err
708 }
709
710
711
712 func (p *ReverseProxy) copyBuffer(dst io.Writer, src io.Reader, buf []byte) (int64, error) {
713 if len(buf) == 0 {
714 buf = make([]byte, 32*1024)
715 }
716 var written int64
717 for {
718 nr, rerr := src.Read(buf)
719 if rerr != nil && rerr != io.EOF && rerr != context.Canceled {
720 p.logf("httputil: ReverseProxy read error during body copy: %v", rerr)
721 }
722 if nr > 0 {
723 nw, werr := dst.Write(buf[:nr])
724 if nw > 0 {
725 written += int64(nw)
726 }
727 if werr != nil {
728 return written, werr
729 }
730 if nr != nw {
731 return written, io.ErrShortWrite
732 }
733 }
734 if rerr != nil {
735 if rerr == io.EOF {
736 rerr = nil
737 }
738 return written, rerr
739 }
740 }
741 }
742
743 func (p *ReverseProxy) logf(format string, args ...any) {
744 if p.ErrorLog != nil {
745 p.ErrorLog.Printf(format, args...)
746 } else {
747 log.Printf(format, args...)
748 }
749 }
750
751 type maxLatencyWriter struct {
752 dst io.Writer
753 flush func() error
754 latency time.Duration
755
756 mu sync.Mutex
757 t *time.Timer
758 flushPending bool
759 }
760
761 func (m *maxLatencyWriter) Write(p []byte) (n int, err error) {
762 m.mu.Lock()
763 defer m.mu.Unlock()
764 n, err = m.dst.Write(p)
765 if m.latency < 0 {
766 m.flush()
767 return
768 }
769 if m.flushPending {
770 return
771 }
772 if m.t == nil {
773 m.t = time.AfterFunc(m.latency, m.delayedFlush)
774 } else {
775 m.t.Reset(m.latency)
776 }
777 m.flushPending = true
778 return
779 }
780
781 func (m *maxLatencyWriter) delayedFlush() {
782 m.mu.Lock()
783 defer m.mu.Unlock()
784 if !m.flushPending {
785 return
786 }
787 m.flush()
788 m.flushPending = false
789 }
790
791 func (m *maxLatencyWriter) stop() {
792 m.mu.Lock()
793 defer m.mu.Unlock()
794 m.flushPending = false
795 if m.t != nil {
796 m.t.Stop()
797 }
798 }
799
800 func upgradeType(h http.Header) string {
801 if !httpguts.HeaderValuesContainsToken(h["Connection"], "Upgrade") {
802 return ""
803 }
804 return h.Get("Upgrade")
805 }
806
807 func (p *ReverseProxy) handleUpgradeResponse(rw http.ResponseWriter, req *http.Request, res *http.Response) {
808 reqUpType := upgradeType(req.Header)
809 resUpType := upgradeType(res.Header)
810 if !ascii.IsPrint(resUpType) {
811 p.getErrorHandler()(rw, req, fmt.Errorf("backend tried to switch to invalid protocol %q", resUpType))
812 return
813 }
814 if !ascii.EqualFold(reqUpType, resUpType) {
815 p.getErrorHandler()(rw, req, fmt.Errorf("backend tried to switch protocol %q when %q was requested", resUpType, reqUpType))
816 return
817 }
818
819 backConn, ok := res.Body.(io.ReadWriteCloser)
820 if !ok {
821 p.getErrorHandler()(rw, req, fmt.Errorf("internal error: 101 switching protocols response with non-writable body"))
822 return
823 }
824
825 rc := http.NewResponseController(rw)
826 conn, brw, hijackErr := rc.Hijack()
827 if errors.Is(hijackErr, http.ErrNotSupported) {
828 p.getErrorHandler()(rw, req, fmt.Errorf("can't switch protocols using non-Hijacker ResponseWriter type %T", rw))
829 return
830 }
831
832 backConnCloseCh := make(chan bool)
833 go func() {
834
835
836 select {
837 case <-req.Context().Done():
838 case <-backConnCloseCh:
839 }
840 backConn.Close()
841 }()
842 defer close(backConnCloseCh)
843
844 if hijackErr != nil {
845 p.getErrorHandler()(rw, req, fmt.Errorf("Hijack failed on protocol switch: %v", hijackErr))
846 return
847 }
848 defer conn.Close()
849
850 copyHeader(rw.Header(), res.Header)
851
852 res.Header = rw.Header()
853 res.Body = nil
854 if err := res.Write(brw); err != nil {
855 p.getErrorHandler()(rw, req, fmt.Errorf("response write: %v", err))
856 return
857 }
858 if err := brw.Flush(); err != nil {
859 p.getErrorHandler()(rw, req, fmt.Errorf("response flush: %v", err))
860 return
861 }
862 errc := make(chan error, 1)
863 spc := switchProtocolCopier{user: conn, backend: backConn}
864 go spc.copyToBackend(errc)
865 go spc.copyFromBackend(errc)
866
867
868
869 err := <-errc
870 if err == nil {
871 err = <-errc
872 }
873 }
874
875 var errCopyDone = errors.New("hijacked connection copy complete")
876
877
878
879 type switchProtocolCopier struct {
880 user, backend io.ReadWriter
881 }
882
883 func (c switchProtocolCopier) copyFromBackend(errc chan<- error) {
884 if _, err := io.Copy(c.user, c.backend); err != nil {
885 errc <- err
886 return
887 }
888
889
890 if wc, ok := c.user.(interface{ CloseWrite() error }); ok {
891 errc <- wc.CloseWrite()
892 return
893 }
894
895 errc <- errCopyDone
896 }
897
898 func (c switchProtocolCopier) copyToBackend(errc chan<- error) {
899 if _, err := io.Copy(c.backend, c.user); err != nil {
900 errc <- err
901 return
902 }
903
904
905 if wc, ok := c.backend.(interface{ CloseWrite() error }); ok {
906 errc <- wc.CloseWrite()
907 return
908 }
909
910 errc <- errCopyDone
911 }
912
913 var urlmaxqueryparams = godebug.New("urlmaxqueryparams")
914
915
916 const defaultMaxParams = 10000
917
918 func cleanQueryParams(s string) string {
919 reencode := func(s string) string {
920 v, _ := url.ParseQuery(s)
921 return v.Encode()
922 }
923 if urlmaxqueryparams.Value() != "" {
924
925 return reencode(s)
926 }
927 if numParams := strings.Count(s, "&") + 1; numParams > defaultMaxParams {
928
929 return reencode(s)
930 }
931 for i := 0; i < len(s); {
932 switch s[i] {
933 case ';':
934 return reencode(s)
935 case '%':
936 if i+2 >= len(s) || !ishex(s[i+1]) || !ishex(s[i+2]) {
937 return reencode(s)
938 }
939 i += 3
940 default:
941 i++
942 }
943 }
944 return s
945 }
946
947 func ishex(c byte) bool {
948 switch {
949 case '0' <= c && c <= '9':
950 return true
951 case 'a' <= c && c <= 'f':
952 return true
953 case 'A' <= c && c <= 'F':
954 return true
955 }
956 return false
957 }
958
View as plain text