knative.dev/pkg@v0.0.0-20260602142205-ac97e43f6622/test/spoof/spoof_test.go (about)

     1  /*
     2  Copyright 2021 The Knative Authors
     3  
     4  Licensed under the Apache License, Version 2.0 (the "License");
     5  you may not use this file except in compliance with the License.
     6  You may obtain a copy of the License at
     7  
     8      http://www.apache.org/licenses/LICENSE-2.0
     9  
    10  Unless required by applicable law or agreed to in writing, software
    11  distributed under the License is distributed on an "AS IS" BASIS,
    12  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
    13  See the License for the specific language governing permissions and
    14  limitations under the License.
    15  */
    16  
    17  // spoof contains logic to make polling HTTP requests against an endpoint with optional host spoofing.
    18  
    19  package spoof
    20  
    21  import (
    22  	"context"
    23  	"errors"
    24  	"net/http"
    25  	"net/url"
    26  	"sync/atomic"
    27  	"testing"
    28  	"time"
    29  )
    30  
    31  var (
    32  	successResponse = &http.Response{
    33  		Status:     "200 ok",
    34  		StatusCode: http.StatusOK,
    35  		Header:     http.Header{},
    36  		Body:       http.NoBody,
    37  	}
    38  	errRetriable    = errors.New("connection reset by peer")
    39  	errNonRetriable = errors.New("foo")
    40  )
    41  
    42  type fakeTransport struct {
    43  	response *http.Response
    44  	err      error
    45  	calls    atomic.Int32
    46  }
    47  
    48  func (ft *fakeTransport) RoundTrip(req *http.Request) (*http.Response, error) {
    49  	call := ft.calls.Add(1)
    50  	if ft.response != nil && ft.err != nil {
    51  		if call == 2 {
    52  			return ft.response, nil
    53  		}
    54  		return nil, ft.err
    55  	} else if ft.response != nil {
    56  		return ft.response, nil
    57  	}
    58  
    59  	return nil, ft.err
    60  }
    61  
    62  func TestSpoofingClient_CheckEndpointState(t *testing.T) {
    63  	tests := []struct {
    64  		name      string
    65  		transport *fakeTransport
    66  		inState   ResponseChecker
    67  		wantErr   bool
    68  		wantCalls int32
    69  	}{{
    70  		name:      "Non matching response doesn't trigger a second check",
    71  		transport: &fakeTransport{response: successResponse},
    72  		inState: func(resp *Response) (done bool, err error) {
    73  			return false, nil
    74  		},
    75  		wantErr:   false,
    76  		wantCalls: 1,
    77  	}, {
    78  		name:      "Error response doesn't trigger a second check",
    79  		transport: &fakeTransport{response: successResponse},
    80  		inState: func(resp *Response) (done bool, err error) {
    81  			return false, errors.New("response error")
    82  		},
    83  		wantErr:   true,
    84  		wantCalls: 1,
    85  	}, {
    86  		name:      "OK response doesn't trigger a second check",
    87  		transport: &fakeTransport{response: successResponse},
    88  		inState: func(resp *Response) (done bool, err error) {
    89  			return true, nil
    90  		},
    91  		wantErr:   false,
    92  		wantCalls: 1,
    93  	}, {
    94  		name:      "Retriable error is retried",
    95  		transport: &fakeTransport{err: errRetriable, response: successResponse},
    96  		inState: func(resp *Response) (done bool, err error) {
    97  			return true, nil
    98  		},
    99  		wantErr:   false,
   100  		wantCalls: 2,
   101  	}, {
   102  		name:      "Nonretriable error is not retried",
   103  		transport: &fakeTransport{err: errNonRetriable, response: successResponse},
   104  		inState: func(resp *Response) (done bool, err error) {
   105  			return true, nil
   106  		},
   107  		wantErr:   true,
   108  		wantCalls: 1,
   109  	}}
   110  	for _, tt := range tests {
   111  		t.Run(tt.name, func(t *testing.T) {
   112  			sc := &SpoofingClient{
   113  				Client:          &http.Client{Transport: tt.transport},
   114  				Logf:            t.Logf,
   115  				RequestInterval: 1,
   116  				RequestTimeout:  time.Second,
   117  			}
   118  			url := &url.URL{
   119  				Host:   "fake.knative.net",
   120  				Scheme: "http",
   121  			}
   122  			_, err := sc.CheckEndpointState(context.TODO(), url, tt.inState, "")
   123  			if (err != nil) != tt.wantErr {
   124  				t.Errorf("SpoofingClient.CheckEndpointState() error = %v, wantErr %v", err, tt.wantErr)
   125  				return
   126  			}
   127  			if got, want := tt.transport.calls.Load(), tt.wantCalls; got != want {
   128  				t.Errorf("Expected Transport to be invoked %d time but got invoked %d", want, got)
   129  			}
   130  		})
   131  	}
   132  }
   133  
   134  func TestSpoofingClient_WaitForEndpointState(t *testing.T) {
   135  	tests := []struct {
   136  		name      string
   137  		transport *fakeTransport
   138  		inState   ResponseChecker
   139  		wantErr   bool
   140  		wantCalls int32
   141  	}{{
   142  		name:      "OK response doesn't trigger a second request",
   143  		transport: &fakeTransport{response: successResponse},
   144  		inState: func(resp *Response) (done bool, err error) {
   145  			return true, nil
   146  		},
   147  		wantErr:   false,
   148  		wantCalls: 1,
   149  	}, {
   150  		name:      "Error response doesn't trigger more requests",
   151  		transport: &fakeTransport{response: successResponse},
   152  		inState: func(resp *Response) (done bool, err error) {
   153  			return false, errors.New("response error")
   154  		},
   155  		wantErr:   true,
   156  		wantCalls: 1,
   157  	}, {
   158  		name:      "Non matching response triggers more requests",
   159  		transport: &fakeTransport{response: successResponse},
   160  		inState: func() ResponseChecker {
   161  			var calls atomic.Int32
   162  			return func(resp *Response) (done bool, err error) {
   163  				val := calls.Add(1)
   164  				// Stop the looping on the third invocation
   165  				return val == 3, nil
   166  			}
   167  		}(),
   168  		wantErr:   false,
   169  		wantCalls: 3,
   170  	}, {
   171  		name:      "Retriable error is retried",
   172  		transport: &fakeTransport{err: errRetriable, response: successResponse},
   173  		inState: func(resp *Response) (done bool, err error) {
   174  			return true, nil
   175  		},
   176  		wantErr:   false,
   177  		wantCalls: 2,
   178  	}, {
   179  		name:      "Nonretriable error is not retried",
   180  		transport: &fakeTransport{err: errNonRetriable, response: successResponse},
   181  		inState: func(resp *Response) (done bool, err error) {
   182  			return true, nil
   183  		},
   184  		wantErr:   true,
   185  		wantCalls: 1,
   186  	}}
   187  	for _, tt := range tests {
   188  		t.Run(tt.name, func(t *testing.T) {
   189  			sc := &SpoofingClient{
   190  				Client:          &http.Client{Transport: tt.transport},
   191  				Logf:            t.Logf,
   192  				RequestInterval: 1,
   193  				RequestTimeout:  time.Second,
   194  			}
   195  			url := &url.URL{
   196  				Host:   "fake.knative.net",
   197  				Scheme: "http",
   198  			}
   199  			_, err := sc.WaitForEndpointState(context.TODO(), url, tt.inState, "")
   200  			if (err != nil) != tt.wantErr {
   201  				t.Errorf("SpoofingClient.CheckEndpointState() error = %v, wantErr %v", err, tt.wantErr)
   202  				return
   203  			}
   204  			if got, want := tt.transport.calls.Load(), tt.wantCalls; got != want {
   205  				t.Errorf("Expected Transport to be invoked %d time but got invoked %d", want, got)
   206  			}
   207  		})
   208  	}
   209  }