Skip to content

Commit e3d594c

Browse files
committed
refactor(node): migrate service proxy from Director to Rewrite
Supports the golang.org/x/time v0.16.0 update in #401. Signed-off-by: Emin Aktas <eminaktas34@gmail.com>
1 parent 79b5ef7 commit e3d594c

2 files changed

Lines changed: 94 additions & 20 deletions

File tree

internal/node/service.go

Lines changed: 8 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -55,32 +55,20 @@ func newReverseProxyHandler(targetURL string) (http.Handler, error) {
5555
if err != nil {
5656
return nil, fmt.Errorf("invalid target URL: %w", err)
5757
}
58-
proxy := httputil.NewSingleHostReverseProxy(u)
59-
if proxy.Rewrite != nil {
60-
originalRewrite := proxy.Rewrite
61-
proxy.Rewrite = func(pr *httputil.ProxyRequest) {
58+
return &httputil.ReverseProxy{
59+
Rewrite: func(pr *httputil.ProxyRequest) {
6260
noTrailingSlash := pr.In.Header.Get(api.HeaderSamNoTrailingSlash) == "true"
63-
originalRewrite(pr)
61+
pr.SetURL(u)
62+
// Preserve the Host behavior of NewSingleHostReverseProxy.
63+
pr.Out.Host = pr.In.Host
6464
pr.Out.Header.Del(api.HeaderSamNoTrailingSlash)
6565
if noTrailingSlash && !strings.HasSuffix(u.Path, "/") && strings.HasSuffix(pr.Out.URL.Path, "/") {
6666
pr.Out.URL.Path = strings.TrimSuffix(pr.Out.URL.Path, "/")
6767
}
68+
pr.SetXForwarded()
6869
logger.Debugf("[ReverseProxy] Forwarding to: %q", pr.Out.URL.String())
69-
}
70-
} else {
71-
originalDirector := proxy.Director
72-
proxy.Director = func(req *http.Request) {
73-
noTrailingSlash := req.Header.Get(api.HeaderSamNoTrailingSlash) == "true"
74-
req.Header.Del(api.HeaderSamNoTrailingSlash)
75-
originalDirector(req)
76-
77-
if noTrailingSlash && !strings.HasSuffix(u.Path, "/") && strings.HasSuffix(req.URL.Path, "/") {
78-
req.URL.Path = strings.TrimSuffix(req.URL.Path, "/")
79-
}
80-
logger.Debugf("[ReverseProxy] Forwarding to: %q", req.URL.String())
81-
}
82-
}
83-
return proxy, nil
70+
},
71+
}, nil
8472
}
8573

8674
func (b *baseService) Info() *api.ServiceInfo { return b.info }

internal/node/service_test.go

Lines changed: 86 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -61,6 +61,92 @@ func TestBaseService_InitURLBackend_BuildsReverseProxy(t *testing.T) {
6161
}
6262
}
6363

64+
func TestNewReverseProxyHandler_RewritesRequests(t *testing.T) {
65+
for _, tc := range []struct {
66+
name string
67+
targetURL string
68+
requestURL string
69+
noTrailingSlash bool
70+
wantURL string
71+
wantProto string
72+
}{
73+
{
74+
name: "joins target path and query",
75+
targetURL: "http://backend.example/base?upstream=one",
76+
requestURL: "https://service.sam/child/?client=two",
77+
wantURL: "http://backend.example/base/child/?upstream=one&client=two",
78+
wantProto: "https",
79+
},
80+
{
81+
name: "marker trims trailing slash",
82+
targetURL: "http://backend.example/base?upstream=one",
83+
requestURL: "http://service.sam/child/?client=two",
84+
noTrailingSlash: true,
85+
wantURL: "http://backend.example/base/child?upstream=one&client=two",
86+
wantProto: "http",
87+
},
88+
{
89+
name: "target trailing slash is retained",
90+
targetURL: "http://backend.example/base/",
91+
requestURL: "http://service.sam/child/",
92+
noTrailingSlash: true,
93+
wantURL: "http://backend.example/base/child/",
94+
wantProto: "http",
95+
},
96+
} {
97+
t.Run(tc.name, func(t *testing.T) {
98+
handler, err := newReverseProxyHandler(tc.targetURL)
99+
if err != nil {
100+
t.Fatalf("newReverseProxyHandler: %v", err)
101+
}
102+
proxy := handler.(*httputil.ReverseProxy)
103+
var forwarded *http.Request
104+
proxy.Transport = roundTripFunc(func(req *http.Request) (*http.Response, error) {
105+
forwarded = req.Clone(req.Context())
106+
return &http.Response{
107+
StatusCode: http.StatusNoContent,
108+
Header: make(http.Header),
109+
Body: http.NoBody,
110+
}, nil
111+
})
112+
req := httptest.NewRequest(http.MethodGet, tc.requestURL, nil)
113+
req.RemoteAddr = "192.0.2.1:12345"
114+
if tc.noTrailingSlash {
115+
req.Header.Set(api.HeaderSamNoTrailingSlash, "true")
116+
}
117+
req.Header.Set("Forwarded", "for=spoofed;host=spoofed;proto=spoofed")
118+
req.Header.Set("X-Forwarded-For", "spoofed")
119+
req.Header.Set("X-Forwarded-Host", "spoofed")
120+
req.Header.Set("X-Forwarded-Proto", "spoofed")
121+
recorder := httptest.NewRecorder()
122+
proxy.ServeHTTP(recorder, req)
123+
if recorder.Code != http.StatusNoContent {
124+
t.Fatalf("status = %d, want %d", recorder.Code, http.StatusNoContent)
125+
}
126+
if forwarded == nil {
127+
t.Fatal("request did not reach the upstream transport")
128+
}
129+
if got := forwarded.URL.String(); got != tc.wantURL {
130+
t.Errorf("upstream URL = %q, want %q", got, tc.wantURL)
131+
}
132+
if forwarded.Host != req.Host {
133+
t.Errorf("upstream Host = %q, want %q", forwarded.Host, req.Host)
134+
}
135+
for name, want := range map[string]string{
136+
api.HeaderSamNoTrailingSlash: "",
137+
"Forwarded": "",
138+
"X-Forwarded-For": "192.0.2.1",
139+
"X-Forwarded-Host": req.Host,
140+
"X-Forwarded-Proto": tc.wantProto,
141+
} {
142+
if got := forwarded.Header.Get(name); got != want {
143+
t.Errorf("upstream %s = %q, want %q", name, got, want)
144+
}
145+
}
146+
})
147+
}
148+
}
149+
64150
func TestBaseService_InitURLBackend_InvalidURL(t *testing.T) {
65151
b := &baseService{
66152
info: &api.ServiceInfo{Type: api.ServiceType_SERVICE_TYPE_INFERENCE, Name: "demo"},

0 commit comments

Comments
 (0)