@@ -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+
64150func TestBaseService_InitURLBackend_InvalidURL (t * testing.T ) {
65151 b := & baseService {
66152 info : & api.ServiceInfo {Type : api .ServiceType_SERVICE_TYPE_INFERENCE , Name : "demo" },
0 commit comments